Skip to content

Commit 01fdf0b

Browse files
author
Utkarsh Simha
committed
fix(quantization): address review feedback on weight FQ calibration change
- Move enable_weight_fake_quant/disable_activation_fake_quant into a new _fake_quant_utils module so they can import FakeQuantizeImplBase and CompressionTargetTensor at module scope instead of via local imports, keeping _utils.py free of dependencies on other quantization modules. - Use basic_config (weight + activation quantization) instead of input_activation_only_config in test_calibration_mode so the weight-FQ branch is actually exercised; only assert scale drift for activation modules, since weight ranges are fixed at prepare time. - Drop the changelog.d entry.
1 parent f0bdcce commit 01fdf0b

6 files changed

Lines changed: 66 additions & 55 deletions

File tree

changelog.d/weight-fq-calibration-mode.changed

Lines changed: 0 additions & 1 deletion
This file was deleted.

src/coreai_opt/quantization/_eager/quantizer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@
3939
apply_weight_axis_defaults_eager as _apply_weight_axis_defaults,
4040
validate_activation_axes as _validate_activation_axes,
4141
)
42-
from coreai_opt.quantization._utils import (
42+
from coreai_opt.quantization._fake_quant_utils import (
4343
disable_activation_fake_quant as _disable_activation_fake_quant,
4444
enable_weight_fake_quant as _enable_weight_fake_quant,
4545
)
Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
# Copyright 2026 Apple Inc.
2+
#
3+
# Use of this source code is governed by a BSD-3-Clause license that can
4+
# be found in the LICENSE file or at https://opensource.org/licenses/BSD-3-Clause
5+
6+
"""Helpers for toggling fake quantization by quantization target."""
7+
8+
import torch
9+
10+
from coreai_opt.config.spec import CompressionTargetTensor
11+
from coreai_opt.quantization.spec.fake_quantize import FakeQuantizeImplBase
12+
13+
14+
def disable_activation_fake_quant(module: torch.nn.Module) -> None:
15+
"""Disable fake quantization on activation FakeQuantize modules only.
16+
17+
Mirrors ``torchao.quantization.pt2e.fake_quantize.disable_fake_quant`` but
18+
skips weight FQ modules. Used by ``calibration_mode`` so activation
19+
observers see the effect of quantized weights when collecting statistics.
20+
21+
Args:
22+
module (torch.nn.Module): Module to (possibly) toggle. No-op for any
23+
module that is not a ``FakeQuantizeImplBase`` whose
24+
``quantization_target`` is ``ACTIVATION``.
25+
"""
26+
if (
27+
isinstance(module, FakeQuantizeImplBase)
28+
and module.quantization_target == CompressionTargetTensor.ACTIVATION
29+
):
30+
module.disable_fake_quant()
31+
32+
33+
def enable_weight_fake_quant(module: torch.nn.Module) -> None:
34+
"""Enable fake quantization on weight FakeQuantize modules only.
35+
36+
Mirrors ``torchao.quantization.pt2e.fake_quantize.enable_fake_quant`` but
37+
skips activation FQ modules. Companion to
38+
:func:`disable_activation_fake_quant` used by ``calibration_mode``.
39+
40+
Args:
41+
module (torch.nn.Module): Module to (possibly) toggle. No-op for any
42+
module that is not a ``FakeQuantizeImplBase`` whose
43+
``quantization_target`` is ``WEIGHT``.
44+
"""
45+
if (
46+
isinstance(module, FakeQuantizeImplBase)
47+
and module.quantization_target == CompressionTargetTensor.WEIGHT
48+
):
49+
module.enable_fake_quant()

src/coreai_opt/quantization/_graph/quantizer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,7 @@
6060
apply_weight_axis_defaults_graph as _apply_weight_axis_defaults,
6161
validate_activation_axes as _validate_activation_axes,
6262
)
63-
from coreai_opt.quantization._utils import (
63+
from coreai_opt.quantization._fake_quant_utils import (
6464
disable_activation_fake_quant as _disable_activation_fake_quant,
6565
enable_weight_fake_quant as _enable_weight_fake_quant,
6666
)

src/coreai_opt/quantization/_utils.py

Lines changed: 0 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -46,51 +46,3 @@ def get_quantization_shapes(
4646
reduced_shape[i] = 1
4747

4848
return original_shape, blockwise_shape, reduced_shape
49-
50-
51-
def disable_activation_fake_quant(module: torch.nn.Module) -> None:
52-
"""Disable fake quantization on activation FakeQuantize modules only.
53-
54-
Mirrors ``torchao.quantization.pt2e.fake_quantize.disable_fake_quant`` but
55-
skips weight FQ modules. Used by ``calibration_mode`` so activation
56-
observers see the effect of quantized weights when collecting statistics.
57-
58-
Args:
59-
module (torch.nn.Module): Module to (possibly) toggle. No-op for any
60-
module that is not a ``FakeQuantizeImplBase`` whose
61-
``quantization_target`` is ``ACTIVATION``.
62-
"""
63-
# Lazy imports break a cycle: spec.fake_quantize imports get_quantization_shapes
64-
# from this module.
65-
from coreai_opt.config.spec import CompressionTargetTensor # noqa: PLC0415
66-
from coreai_opt.quantization.spec.fake_quantize import FakeQuantizeImplBase # noqa: PLC0415
67-
68-
if (
69-
isinstance(module, FakeQuantizeImplBase)
70-
and module.quantization_target == CompressionTargetTensor.ACTIVATION
71-
):
72-
module.disable_fake_quant()
73-
74-
75-
def enable_weight_fake_quant(module: torch.nn.Module) -> None:
76-
"""Enable fake quantization on weight FakeQuantize modules only.
77-
78-
Mirrors ``torchao.quantization.pt2e.fake_quantize.enable_fake_quant`` but
79-
skips activation FQ modules. Companion to
80-
:func:`disable_activation_fake_quant` used by ``calibration_mode``.
81-
82-
Args:
83-
module (torch.nn.Module): Module to (possibly) toggle. No-op for any
84-
module that is not a ``FakeQuantizeImplBase`` whose
85-
``quantization_target`` is ``WEIGHT``.
86-
"""
87-
# Lazy imports break a cycle: spec.fake_quantize imports get_quantization_shapes
88-
# from this module.
89-
from coreai_opt.config.spec import CompressionTargetTensor # noqa: PLC0415
90-
from coreai_opt.quantization.spec.fake_quantize import FakeQuantizeImplBase # noqa: PLC0415
91-
92-
if (
93-
isinstance(module, FakeQuantizeImplBase)
94-
and module.quantization_target == CompressionTargetTensor.WEIGHT
95-
):
96-
module.enable_fake_quant()

tests/quantization/test_eager_quant.py

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -530,23 +530,30 @@ def test_finalize_state_dict_safetensors_roundtrip(self, basic_config, tmp_path)
530530

531531
assert torch.equal(out_before_roundtrip, out_after_roundtrip)
532532

533-
def test_calibration_mode(self, simple_model, input_activation_only_config, example_input):
533+
def test_calibration_mode(self, simple_model, basic_config, example_input):
534534
"""
535535
Test that calibration mode works as expected, and scales are getting updated
536536
"""
537-
quantizer = Quantizer(simple_model, input_activation_only_config)
537+
quantizer = Quantizer(simple_model, basic_config)
538538
simple_model.eval()
539539
prepared_model = quantizer.prepare((example_input,))
540540

541541
fake_quant_modules = [
542542
m for m in prepared_model.modules() if isinstance(m, FakeQuantizeImplBase)
543543
]
544+
activation_fake_quant_modules = [
545+
m
546+
for m in fake_quant_modules
547+
if m.quantization_target == CompressionTargetTensor.ACTIVATION
548+
]
544549

545550
for module in fake_quant_modules:
546551
assert module.observer_enabled.item() == 0
547552
assert module.fake_quant_enabled.item() == 1
548553

549-
pre_calibration_scales = [mod.calculate_qparams()[0].clone() for mod in fake_quant_modules]
554+
pre_calibration_scales = [
555+
mod.calculate_qparams()[0].clone() for mod in activation_fake_quant_modules
556+
]
550557

551558
with quantizer.calibration_mode():
552559
simple_model.eval()
@@ -563,8 +570,12 @@ def test_calibration_mode(self, simple_model, input_activation_only_config, exam
563570
)
564571
assert module.fake_quant_enabled.item() == expected_fq
565572

566-
post_calibration_scales = [mod.calculate_qparams()[0].clone() for mod in fake_quant_modules]
573+
post_calibration_scales = [
574+
mod.calculate_qparams()[0].clone() for mod in activation_fake_quant_modules
575+
]
567576

577+
# Only activation scales are expected to move here: weight ranges are
578+
# fixed at prepare time and don't depend on calibration data.
568579
for pre_scale, post_scale in zip(
569580
pre_calibration_scales, post_calibration_scales, strict=True
570581
):

0 commit comments

Comments
 (0)