Skip to content

Commit e9e8e7b

Browse files
committed
Address review comments
1 parent 291217d commit e9e8e7b

3 files changed

Lines changed: 22 additions & 62 deletions

File tree

src/coreai_opt/quantization/_graph/_annotation_utils.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -158,10 +158,13 @@ class OpsListPattern:
158158
# Dictionary mapping ops with known output bounds to (qscheme, float_range).
159159
# float_range elements may be None to leave that side data-driven.
160160
_fixed_q_params_ops = {
161+
# tanh: bounded to [-1, 1]
161162
torch.ops.aten.tanh.default: (QuantizationScheme.SYMMETRIC, (-1.0, 1.0)),
162163
torch.ops.aten.tanh_.default: (QuantizationScheme.SYMMETRIC, (-1.0, 1.0)),
164+
# sigmoid: bounded to [0, 1]
163165
torch.ops.aten.sigmoid.default: (QuantizationScheme.ASYMMETRIC, (0.0, 1.0)),
164166
torch.ops.aten.sigmoid_.default: (QuantizationScheme.ASYMMETRIC, (0.0, 1.0)),
167+
# hardsigmoid: bounded to [0, 1]
165168
torch.ops.aten.hardsigmoid.default: (QuantizationScheme.ASYMMETRIC, (0.0, 1.0)),
166169
torch.ops.aten.hardsigmoid_.default: (QuantizationScheme.ASYMMETRIC, (0.0, 1.0)),
167170
# relu: always >= 0, upper bound is data-driven
@@ -271,7 +274,9 @@ def _propagate_adjusted_spec_to_child_nodes(
271274
if float_range is not None:
272275
kwargs["float_range"] = float_range
273276
if kwargs:
274-
ctr = QuantizationComponentFactory.update_partial_qparams_calculator(ctr, **kwargs)
277+
ctr = QuantizationComponentFactory.reconstruct_partial_qparams_calculator(
278+
ctr, **kwargs
279+
)
275280

276281
# qscheme in TorchAOQuantizationSpec is not read by coreai-opt later on so we omit it.
277282
# Only the qscheme contained within observer_or_fake_quant_ctr matters.
@@ -319,7 +324,7 @@ def adjust_output_qspec_for_qscheme_and_propagate(
319324
else:
320325
return
321326

322-
ctr = QuantizationComponentFactory.update_partial_qparams_calculator(
327+
ctr = QuantizationComponentFactory.reconstruct_partial_qparams_calculator(
323328
qspec.observer_or_fake_quant_ctr, qscheme=qscheme, float_range=float_range
324329
)
325330

src/coreai_opt/quantization/spec/factory.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -213,7 +213,7 @@ def create_fake_quantizer(
213213
return spec.fake_quantize_cls(**common_args, **extra_args)
214214

215215
@classmethod
216-
def update_partial_qparams_calculator(
216+
def reconstruct_partial_qparams_calculator(
217217
cls,
218218
partial_ctr: _PartialConstructor,
219219
**kwargs: Any,

tests/quantization/test_factory.py

Lines changed: 14 additions & 59 deletions
Original file line numberDiff line numberDiff line change
@@ -604,9 +604,9 @@ def test_partial_vs_direct_instantiation(self):
604604
assert fq_partial_out.dtype == x.dtype
605605
assert fq_direct_out.dtype == x.dtype
606606

607-
def test_update_partial_qparams_calculator_single_attr(self):
608-
"""
609-
update_partial_qparams_calculator overrides one attribute on each constructed calculator.
607+
def test_reconstruct_partial_qparams_calculator(self):
608+
"""reconstruct_partial_qparams_calculator overrides attributes without mutating the original
609+
partial.
610610
"""
611611
spec = QuantizationSpec(
612612
dtype=torch.int8,
@@ -619,78 +619,33 @@ def test_update_partial_qparams_calculator_single_attr(self):
619619
spec, quantization_target=CompressionTargetTensor.WEIGHT
620620
)
621621

622-
updated = QuantizationComponentFactory.update_partial_qparams_calculator(
622+
# Overriding a single attribute updates the constructed calculator.
623+
updated_single = QuantizationComponentFactory.reconstruct_partial_qparams_calculator(
623624
partial, qscheme=QuantizationScheme.ASYMMETRIC
624625
)
625-
626-
fq = updated()
626+
fq = updated_single()
627627
assert fq.qparams_calculator.qscheme == QuantizationScheme.ASYMMETRIC
628628
# qscheme property on the fake quantizer delegates to the calculator
629629
assert fq.qscheme == QuantizationScheme.ASYMMETRIC
630630

631-
def test_update_partial_qparams_calculator_multiple_attrs(self):
632-
"""update_partial_qparams_calculator overrides multiple attributes in one call."""
633-
spec = QuantizationSpec(
634-
dtype=torch.int8,
635-
qscheme="symmetric",
636-
granularity=PerTensorGranularity(),
637-
qparam_calculator_cls=StaticQParamsCalculator,
638-
range_calculator_cls=MinMaxRangeCalculator,
639-
)
640-
partial = QuantizationComponentFactory.create_fake_quantizer_partial(
641-
spec, quantization_target=CompressionTargetTensor.WEIGHT
642-
)
643-
644-
updated = QuantizationComponentFactory.update_partial_qparams_calculator(
631+
# Overriding multiple attributes in one call updates all of them.
632+
updated_multiple = QuantizationComponentFactory.reconstruct_partial_qparams_calculator(
645633
partial,
646634
qscheme=QuantizationScheme.ASYMMETRIC,
647635
float_range=(0.0, 1.0),
648636
)
649-
650-
fq = updated()
637+
fq = updated_multiple()
651638
assert fq.qparams_calculator.qscheme == QuantizationScheme.ASYMMETRIC
652639
assert fq.qparams_calculator.float_range == (0.0, 1.0)
653640

654-
def test_update_partial_qparams_calculator_does_not_mutate_original(self):
655-
"""update_partial_qparams_calculator returns a new partial; the original is unchanged."""
656-
spec = QuantizationSpec(
657-
dtype=torch.int8,
658-
qscheme="symmetric",
659-
granularity=PerTensorGranularity(),
660-
qparam_calculator_cls=StaticQParamsCalculator,
661-
range_calculator_cls=MinMaxRangeCalculator,
662-
)
663-
partial = QuantizationComponentFactory.create_fake_quantizer_partial(
664-
spec, quantization_target=CompressionTargetTensor.WEIGHT
665-
)
666-
667-
_ = QuantizationComponentFactory.update_partial_qparams_calculator(
668-
partial, qscheme=QuantizationScheme.ASYMMETRIC
669-
)
670-
671-
# Original partial still produces calculators with the original qscheme.
641+
# The original partial is unaffected by either override.
672642
fq_original = partial()
673643
assert fq_original.qparams_calculator.qscheme == QuantizationScheme.SYMMETRIC
674644

675-
def test_update_partial_qparams_calculator_independent_instances(self):
676-
"""Each call to an updated partial creates an independent calculator with the override."""
677-
spec = QuantizationSpec(
678-
dtype=torch.int8,
679-
qscheme="symmetric",
680-
granularity=PerTensorGranularity(),
681-
qparam_calculator_cls=MovingAverageQParamsCalculator,
682-
range_calculator_cls=MinMaxRangeCalculator,
683-
)
684-
partial = QuantizationComponentFactory.create_fake_quantizer_partial(
685-
spec, quantization_target=CompressionTargetTensor.ACTIVATION
686-
)
687-
updated = QuantizationComponentFactory.update_partial_qparams_calculator(
688-
partial, float_range=(0.0, 1.0)
689-
)
690-
691-
fq1 = updated()
692-
fq2 = updated()
693-
645+
# Each call to an updated partial creates an independent calculator with the override
646+
# applied.
647+
fq1 = updated_multiple()
648+
fq2 = updated_multiple()
694649
assert id(fq1.qparams_calculator) != id(fq2.qparams_calculator)
695650
assert fq1.qparams_calculator.float_range == (0.0, 1.0)
696651
assert fq2.qparams_calculator.float_range == (0.0, 1.0)

0 commit comments

Comments
 (0)