Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions src/nncf/common/quantization/quantizer_setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,8 +72,19 @@ def from_state(cls, state: dict[str, Any]) -> "QuantizationInsertionPointBase":
return cls(**state)


class WQIPointStateNames:
INPUT_PORT_ID = "input_port_id"
TARGET_NODE_NAME = "target_node_name"


@CommonStatefulClassesRegistry.register()
class WeightQuantizationInsertionPoint(QuantizationInsertionPointBase):
_state_names = WQIPointStateNames # type: ignore[assignment]

def __init__(self, target_node_name: NNCFNodeName, input_port_id: int | None = None):
super().__init__(target_node_name)
self.input_port_id = input_port_id

def __eq__(self, other: object) -> bool:
return isinstance(other, WeightQuantizationInsertionPoint) and self.target_node_name == other.target_node_name

Expand All @@ -83,6 +94,12 @@ def __str__(self) -> str:
def __hash__(self) -> int:
return hash(str(self))

def get_state(self) -> dict[str, Any]:
state: dict[str, Any] = {WQIPointStateNames.TARGET_NODE_NAME: self.target_node_name}
if self.input_port_id is not None:
state[WQIPointStateNames.INPUT_PORT_ID] = self.input_port_id
return state


class AQIPointStateNames:
INPUT_PORT_ID = "input_port_id"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,10 @@
from nncf.experimental.quantization.quantizer import Quantizer
from nncf.experimental.torch.fx.nncf_graph_builder import GraphConverter
from nncf.experimental.torch.fx.node_utils import get_node_args
from nncf.experimental.torch.fx.node_utils import get_tensor_constant_from_node
from nncf.quantization.algorithms.weight_compression.config import WeightCompressionParameters
from nncf.tensor.definitions import TensorDataType
from nncf.torch.graph import operator_metatypes as om

EdgeOrNode = tuple[torch.fx.Node, torch.fx.Node]

Expand Down Expand Up @@ -88,12 +90,19 @@ def _get_quantization_points(
to_node = to_nodes[0]
if from_node.op == "get_attr":
_, metatype = GraphConverter.get_node_type_and_metatype(to_node, annotated_model)
# Check that the constant is placed on the actual weight port, as it is possible for
# activations to be a constant as well.
if get_node_args(to_node).index(from_node) in metatype.weight_port_ids:
input_port_id = get_node_args(to_node).index(from_node)

if input_port_id in metatype.weight_port_ids:
qip = WeightQuantizationInsertionPoint(to_node.name)
return [SingleConfigQuantizationPoint(qip, qconfig, [x.name for x in to_nodes])]

# Mul has no fixed weight port because either operand may be a learned parameter.
is_mul_parameter = metatype is om.PTMulMetatype and isinstance(
get_tensor_constant_from_node(from_node, annotated_model), torch.nn.Parameter
)
if is_mul_parameter:
qip = WeightQuantizationInsertionPoint(to_node.name, input_port_id)
return [SingleConfigQuantizationPoint(qip, qconfig, [x.name for x in to_nodes])]
if len(from_node.users) == len(to_nodes):
qip = ActivationQuantizationInsertionPoint(from_node.name)
return [SingleConfigQuantizationPoint(qip, qconfig, [x.name for x in to_nodes])]
Expand Down
7 changes: 6 additions & 1 deletion src/nncf/quantization/algorithms/min_max/algorithm.py
Original file line number Diff line number Diff line change
Expand Up @@ -764,7 +764,12 @@ def _get_weight_quantization_target_points(
weight_quantization_target_points = []
node_name = quantization_point.insertion_point.target_node_name
node = nncf_graph.get_node_by_name(node_name)
weights_port_ids = self._backend_entity.get_weight_tensor_port_ids(node, nncf_graph)
input_port_id = quantization_point.insertion_point.input_port_id
weights_port_ids = (
[input_port_id]
if input_port_id is not None
else self._backend_entity.get_weight_tensor_port_ids(node, nncf_graph)
)
for port_id in weights_port_ids:
weight_quantization_target_points.append(
self._backend_entity.target_point(TargetType.OPERATION_WITH_WEIGHTS, node_name, port_id)
Expand Down
208 changes: 208 additions & 0 deletions tests/executorch/test_ptq.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
# limitations under the License.

import json
import logging
from dataclasses import dataclass
from functools import partial
from typing import Any, Callable
Expand All @@ -35,6 +36,7 @@
import nncf
from nncf.common.graph import NNCFGraph
from nncf.common.utils.os import safe_open
from nncf.experimental.quantization.algorithms.post_training.algorithm import ExperimentalPostTrainingQuantization
from nncf.experimental.torch.fx import quantize_pt2e
from nncf.experimental.torch.fx.nncf_graph_builder import GraphConverter
from nncf.experimental.torch.fx.node_utils import get_graph_node_by_name
Expand Down Expand Up @@ -402,6 +404,212 @@ def validate(self, model):
return


class OneInputAnnotationQuantizer(TorchAOQuantizer):
def __init__(self, node_name: str, input_node_name: str, qspec: TorchAOQuantizationSpec):
self._node_name = node_name
self._input_node_name = input_node_name
self._qspec = qspec

def annotate(self, model: torch.fx.GraphModule):
target_node = get_graph_node_by_name(model.graph, self._node_name)
input_node = get_graph_node_by_name(model.graph, self._input_node_name)
target_node.meta["quantization_annotation"] = QuantizationAnnotation(input_qspec_map={input_node: self._qspec})
return model

def validate(self, model):
return


class ActivationTimesWeightModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.ones(3, 1, 1))

def forward(self, x):
return x * self.weight


class WeightTimesActivationModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.ones(3, 1, 1))

def forward(self, x):
return self.weight * x


class ActivationTimesBufferModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.register_buffer("scale", torch.ones(3, 1, 1))

def forward(self, x):
return x * self.scale


def _get_int8_per_tensor_symmetric_qspec() -> TorchAOQuantizationSpec:
return TorchAOQuantizationSpec(
dtype=torch.int8,
observer_or_fake_quant_ctr=None,
qscheme=torch.per_tensor_symmetric,
)


def _get_mul_input_quantization_point_class(model: torch.nn.Module, input_node_name: str) -> str:
example_input = torch.ones(2, 3, 8, 8)
fx_model = get_torch_fx_model(model.eval(), example_input)
input_node = get_graph_node_by_name(fx_model.graph, input_node_name)
annotation = QuantizationAnnotation(input_qspec_map={input_node: _get_int8_per_tensor_symmetric_qspec()})
adapter = TorchAOQuantizerAdapter(OneNodeAnnotationQuantizer("mul", annotation))
setup = adapter.get_quantization_setup(fx_model, GraphConverter.create_nncf_graph(fx_model))
quantization_point = next(iter(setup.get_state()["quantization_points"].values()))
return quantization_point["qip_class"]


@pytest.mark.parametrize(
"model_cls",
[
ActivationTimesWeightModel,
WeightTimesActivationModel,
],
)
def test_torch_ao_adapter_detects_parameter_as_mul_weight(model_cls):
assert _get_mul_input_quantization_point_class(model_cls(), "weight") == "WeightQuantizationInsertionPoint"


def test_torch_ao_adapter_does_not_detect_buffer_as_mul_weight():
assert (
_get_mul_input_quantization_point_class(ActivationTimesBufferModel(), "scale")
== "ActivationQuantizationInsertionPoint"
)


@pytest.mark.parametrize(
("model_cls", "input_node_name"),
[
(ActivationTimesWeightModel, "weight"),
(WeightTimesActivationModel, "weight"),
(ActivationTimesBufferModel, "scale"),
],
)
def test_quantize_pt2e_mul_get_attr_input(model_cls, input_node_name):
example_input = torch.ones(2, 3, 8, 8)
fx_model = get_torch_fx_model(model_cls().eval(), example_input)

quantizer = OneInputAnnotationQuantizer(
"mul",
input_node_name,
_get_int8_per_tensor_symmetric_qspec(),
)

data_loader = torch.utils.data.DataLoader(
torch.ones(2, 3, 8, 8),
batch_size=2,
)
calibration_dataset = nncf.Dataset(data_loader, lambda x: x.to("cpu"))

quantized_model = quantize_pt2e(
fx_model,
quantizer,
calibration_dataset=calibration_dataset,
subset_size=1,
fast_bias_correction=None,
fold_quantize=False,
do_copy=True,
)

input_node = get_graph_node_by_name(quantized_model.graph, input_node_name)
quantize_node = next(iter(input_node.users))

assert quantize_node.target == torch.ops.quantized_decomposed.quantize_per_tensor.default

dequantize_node = next(iter(quantize_node.users))
assert dequantize_node.target == torch.ops.quantized_decomposed.dequantize_per_tensor.default

mul_node = get_graph_node_by_name(quantized_model.graph, "mul")
assert dequantize_node in mul_node.args


@pytest.mark.parametrize(
"model_cls",
[
ActivationTimesWeightModel,
WeightTimesActivationModel,
],
)
def test_quantize_pt2e_mul_parameter_range_with_batchwise_statistics(
model_cls: type[torch.nn.Module],
) -> None:
model = model_cls().eval()
with torch.no_grad():
model.weight.copy_(torch.tensor([0.25, -0.5, 1.0]).reshape(3, 1, 1))

example_input = torch.ones(2, 3, 8, 8)
fx_model = get_torch_fx_model(model, example_input)
quantizer = OneInputAnnotationQuantizer(
"mul",
"weight",
_get_int8_per_tensor_symmetric_qspec(),
)

data_loader = torch.utils.data.DataLoader(
torch.ones(2, 3, 8, 8),
batch_size=2,
)
calibration_dataset = nncf.Dataset(data_loader, lambda x: x.to("cpu"))

quantized_model = quantize_pt2e(
fx_model,
quantizer,
calibration_dataset=calibration_dataset,
subset_size=1,
fast_bias_correction=None,
fold_quantize=False,
do_copy=True,
)

weight_node = get_graph_node_by_name(quantized_model.graph, "weight")
quantize_node = next(iter(weight_node.users))

assert quantize_node.target == torch.ops.quantized_decomposed.quantize_per_tensor.default

scale, zero_point, _, quant_max = quantize_node.args[1:5]
positive_range_bound = (quant_max - zero_point) * scale

# The expected weight range is based on absmax=1.0. If axis 0 were treated
# as a batch axis, the bound would instead be mean(abs(weight)) ~= 0.5833.
assert positive_range_bound == pytest.approx(1.0, rel=1e-5)


def test_quantize_pt2e_custom_quantizer_with_batchwise_statistics(caplog, mocker):
example_input = torch.ones(2, 3, 3, 3)
fx_model = get_torch_fx_model(LinearModel(torch.ones(3, 3)), example_input)
annotation = QuantizationAnnotation(output_qspec=_get_int8_per_tensor_symmetric_qspec())
quantizer = OneNodeAnnotationQuantizer("linear", annotation)
data_loader = torch.utils.data.DataLoader(torch.ones(2, 3, 3, 3), batch_size=2)
calibration_dataset = nncf.Dataset(data_loader, lambda x: x.to("cpu"))
ptq_init_spy = mocker.spy(ExperimentalPostTrainingQuantization, "__init__")
caplog.set_level(logging.INFO, logger="nncf")

quantized_model = quantize_pt2e(
fx_model,
quantizer,
calibration_dataset=calibration_dataset,
subset_size=1,
fast_bias_correction=None,
fold_quantize=False,
do_copy=True,
)

assert calibration_dataset.get_batch_size() == 2
assert ptq_init_spy.call_args.kwargs["batchwise_statistics"] is True
assert "The model has no operations to apply quantization" not in caplog.text

node_targets = {node.target for node in quantized_model.graph.nodes}
assert torch.ops.quantized_decomposed.quantize_per_tensor.default in node_targets
assert torch.ops.quantized_decomposed.dequantize_per_tensor.default in node_targets


REF_NONE_Q_MIN_Q_MAX_SETUP = {
"quantization_points": {
0: {
Expand Down
Loading