diff --git a/modelopt/torch/quantization/model_calib.py b/modelopt/torch/quantization/model_calib.py index d4266b289b9..1ef6f05b884 100644 --- a/modelopt/torch/quantization/model_calib.py +++ b/modelopt/torch/quantization/model_calib.py @@ -20,6 +20,7 @@ import time import warnings from collections.abc import Callable, Mapping, Sequence +from contextlib import ExitStack from functools import partial from typing import Any, TypeAlias @@ -34,6 +35,7 @@ from modelopt.torch.quantization.utils.layerwise_calib import ( LayerActivationCollector, _CheckpointState, + _hide_modules_from_traversal, _reconcile_export_with_resume, ) from modelopt.torch.utils import print_rank_0, warn_rank_0 @@ -2084,6 +2086,20 @@ def layerwise_calibrate( "Layerwise calibration requires a model with identifiable transformer layers." ) + decoder_owned_ids = {id(module) for layer in transformer_layers for module in layer.modules()} + has_enabled_outside_quantizer = any( + isinstance(module, TensorQuantizer) + and module.is_enabled + and id(module) not in decoder_owned_ids + for module in model.modules() + ) + + if export_dir is not None and has_enabled_outside_quantizer: + raise ValueError( + "Layerwise export does not support enabled quantizers outside transformer layers. " + "Calibrate without export_dir, then export the completed model separately." + ) + num_layers = len(transformer_layers) print_rank_0(f"Layerwise calibration: Found {num_layers} transformer layers") @@ -2205,6 +2221,27 @@ def _layer_forward_loop(m, _inputs=layer_inputs): if ckpt: ckpt.full_restore(transformer_layers, model) + if has_enabled_outside_quantizer: + if any(device == "disk" for device in getattr(model, "hf_device_map", {}).values()): + warn_rank_0( + "Layerwise calibration found enabled quantizers outside transformer layers. " + "The required full-model calibration pass may be slow because disk-offloaded " + "decoder weights can be streamed for every batch." + ) + + with _hide_modules_from_traversal(model, transformer_layers): + if qdq_from_prev: + calib_func(model, forward_loop, **calib_kwargs) + else: + with ExitStack() as stack: + for layer in transformer_layers: + stack.enter_context( + set_quantizer_by_cfg_context( + layer, [{"quantizer_name": "*", "enable": False}] + ) + ) + calib_func(model, forward_loop, **calib_kwargs) + if exporter is not None: exporter.finalize() print_rank_0(f"Layerwise export: wrote quantized checkpoint to {export_dir}") diff --git a/modelopt/torch/quantization/utils/layerwise_calib.py b/modelopt/torch/quantization/utils/layerwise_calib.py index e04f83d380b..542445e3526 100644 --- a/modelopt/torch/quantization/utils/layerwise_calib.py +++ b/modelopt/torch/quantization/utils/layerwise_calib.py @@ -27,6 +27,7 @@ import os import shutil from collections import deque +from contextlib import contextmanager from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any @@ -42,7 +43,7 @@ ) if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Sequence from modelopt.torch.opt.searcher import ForwardLoop @@ -101,6 +102,48 @@ def forward(self, *args, **kwargs): ) +class _ForwardOnlyLayer(nn.Module): + """Hide a layer from module traversal while preserving its forward execution.""" + + _PROXY_BLOCKLIST = _SkipLayer._PROXY_BLOCKLIST + + def __init__(self, original: nn.Module): + super().__init__() + object.__setattr__(self, "_original", original) + + def __getattr__(self, name: str): + try: + return super().__getattr__(name) + except AttributeError: + if name in self._PROXY_BLOCKLIST: + raise + return getattr(object.__getattribute__(self, "_original"), name) + + def forward(self, *args, **kwargs): + return self._original(*args, **kwargs) + + +@contextmanager +def _hide_modules_from_traversal(model: nn.Module, modules: Sequence[nn.Module]): + """Temporarily hide registered modules while retaining their forward behavior.""" + target_ids = {id(module) for module in modules} + slots = [ + (parent, child_name, child) + for parent in tuple(model.modules()) + for child_name, child in tuple(parent._modules.items()) + if child is not None and id(child) in target_ids + ] + proxies = {id(child): _ForwardOnlyLayer(child) for _, _, child in slots} + + try: + for parent, child_name, child in slots: + parent._modules[child_name] = proxies[id(child)] + yield + finally: + for parent, child_name, child in slots: + parent._modules[child_name] = child + + class LayerActivationCollector: """Collects layer activations for layerwise (layer-by-layer) calibration. diff --git a/tests/unit/torch/quantization/test_layerwise_calibrate.py b/tests/unit/torch/quantization/test_layerwise_calibrate.py index bda8c6029b1..a9b4ab00471 100644 --- a/tests/unit/torch/quantization/test_layerwise_calibrate.py +++ b/tests/unit/torch/quantization/test_layerwise_calibrate.py @@ -18,6 +18,7 @@ import copy import json from collections import deque +from contextlib import nullcontext import pytest import torch @@ -26,7 +27,12 @@ import modelopt.torch.quantization as mtq from modelopt.torch.quantization.model_calib import layerwise_calibrate from modelopt.torch.quantization.nn import TensorQuantizer -from modelopt.torch.quantization.utils.layerwise_calib import LayerActivationCollector, _SkipLayer +from modelopt.torch.quantization.utils.layerwise_calib import ( + LayerActivationCollector, + _ForwardOnlyLayer, + _hide_modules_from_traversal, + _SkipLayer, +) class _DecoderBlock(nn.Module): @@ -63,6 +69,15 @@ def forward(self, x, **kwargs): return x +class _TransformerWithLMHead(_SimpleTransformerModel): + def __init__(self, n_layers=3, dim=16): + super().__init__(n_layers=n_layers, dim=dim) + self.lm_head = nn.Linear(dim, 32, bias=False) + + def forward(self, x, **kwargs): + return self.lm_head(super().forward(x, **kwargs)) + + class _FlatMLP(nn.Module): """No decoder-layer structure -- should be rejected by layerwise_calibrate.""" @@ -74,6 +89,48 @@ def forward(self, x): return self.net(x) +class _TrackingQuantizer(TensorQuantizer): + """Deterministic quantizer used to distinguish QDQ and FP activations.""" + + def __init__(self): + super().__init__(amax=3.0) + self.calls = 0 + + def forward(self, x): + self.calls += 1 + return x / 2 if self.is_enabled else x + + +class _QuantizedLayer(nn.Module): + def __init__(self): + super().__init__() + self.quantizer = _TrackingQuantizer() + + def forward(self, x): + return self.quantizer(x) + + +class _QuantizedTail(nn.Module): + def __init__(self): + super().__init__() + self.quantizer = _TrackingQuantizer() + self.inputs = [] + + def forward(self, x): + self.inputs.append(x.detach().clone()) + return self.quantizer(x) + + +class _ModelWithQuantizedTail(nn.Module): + def __init__(self, with_tail=True): + super().__init__() + self.layers = nn.ModuleList([_QuantizedLayer()]) + self.tail = _QuantizedTail() if with_tail else nn.Identity() + + def forward(self, x): + return self.tail(self.layers[0](x)) + + class _SimpleTwoLayerModel(nn.Module): """Minimal model with explicit layers for activation-collection tests.""" @@ -243,6 +300,122 @@ def test_layerwise_calib_empty_forward_loop_raises(monkeypatch): ) +@pytest.mark.parametrize("raises", [False, True]) +def test_hide_modules_from_traversal_restores_aliases(raises): + model = _ModelWithQuantizedTail() + original = model.layers[0] + model.layer_alias = original + + error_context = pytest.raises(RuntimeError, match="injected") if raises else nullcontext() + with error_context, _hide_modules_from_traversal(model, [original]): + assert isinstance(model.layers[0], _ForwardOnlyLayer) + assert model.layer_alias is model.layers[0] + assert original not in model.modules() + assert original.quantizer not in model.modules() + if raises: + raise RuntimeError("injected") + + assert model.layers[0] is original + assert model.layer_alias is original + + +@pytest.mark.parametrize( + ("qdq_from_prev", "expected_tail_input"), + [(True, 1.0), (False, 2.0)], +) +def test_layerwise_calibrates_only_outside_quantizers_with_full_model_forward( + monkeypatch, qdq_from_prev, expected_tail_input +): + monkeypatch.setattr( + LayerActivationCollector, + "_decoder_layer_support", + [(lambda m: hasattr(m, "layers"), lambda m: list(m.layers))], + ) + model = _ModelWithQuantizedTail() + decoder_quantizer = model.layers[0].quantizer + calibrated_quantizers = [] + calibrated_targets = [] + decoder_amax_before_extra_pass = [] + + def calib_func(target, target_forward_loop): + calibrated_targets.append(target) + if target is model: + decoder_amax_before_extra_pass.append(decoder_quantizer._amax.clone()) + calibrated_quantizers.append( + {id(module) for module in target.modules() if isinstance(module, TensorQuantizer)} + ) + target_forward_loop(target) + + layerwise_calibrate( + model, + lambda m: m(torch.tensor([2.0])), + calib_func, + get_qdq_activations_from_prev_layer=qdq_from_prev, + ) + + assert calibrated_targets == [model.layers[0], model] + assert calibrated_quantizers == [{id(decoder_quantizer)}, {id(model.tail.quantizer)}] + torch.testing.assert_close(model.tail.inputs[-1], torch.tensor([expected_tail_input])) + torch.testing.assert_close(decoder_quantizer._amax, decoder_amax_before_extra_pass[0]) + assert decoder_quantizer.calls == (2 if qdq_from_prev else 1) + assert decoder_quantizer.is_enabled + + +def test_layerwise_skips_full_model_pass_without_outside_quantizer(monkeypatch): + _register_test_discoverer(monkeypatch) + model = _ModelWithQuantizedTail(with_tail=False) + calibrated_targets = [] + + def calib_func(target, target_forward_loop): + calibrated_targets.append(target) + target_forward_loop(target) + + layerwise_calibrate(model, lambda m: m(torch.tensor([2.0])), calib_func) + + assert calibrated_targets == [model.layers[0]] + + +@pytest.mark.parametrize( + ("device_map", "with_tail", "warns"), + [ + ({"layers.0": "disk"}, True, True), + ({"layers.0": "cpu"}, True, False), + ({}, True, False), + ({"layers.0": "disk"}, False, False), + ], +) +def test_layerwise_disk_offload_warning_gating(monkeypatch, device_map, with_tail, warns): + _register_test_discoverer(monkeypatch) + model = _ModelWithQuantizedTail(with_tail=with_tail) + model.hf_device_map = device_map + warnings = [] + monkeypatch.setattr("modelopt.torch.quantization.model_calib.warn_rank_0", warnings.append) + + def calib_func(target, target_forward_loop): + target_forward_loop(target) + + layerwise_calibrate(model, lambda m: m(torch.tensor([2.0])), calib_func) + + assert bool(warnings) is warns + if warns: + assert "disk-offloaded" in warnings[0] + + +def test_layerwise_export_rejects_enabled_outside_quantizer(monkeypatch, tmp_path): + _register_test_discoverer(monkeypatch) + model = _ModelWithQuantizedTail() + + with pytest.raises(ValueError, match="outside transformer layers"): + layerwise_calibrate( + model, + lambda m: m(torch.tensor([2.0])), + lambda *_args, **_kwargs: None, + export_dir=str(tmp_path / "export"), + ) + + assert not (tmp_path / "export").exists() + + # --------------------------------------------------------------------------- # Skip / run / capture path verification tests # --------------------------------------------------------------------------- @@ -725,6 +898,25 @@ def forward_loop(m): model(calib_data[0]) +def test_mtq_quantize_layerwise_calibrates_lm_head(monkeypatch): + _register_test_discoverer(monkeypatch) + config = _int8_cfg_with_algorithm( + { + "method": "max", + "layerwise": {"enable": True, "get_qdq_activations_from_prev_layer": True}, + } + ) + config["quant_cfg"].append({"quantizer_name": "*lm_head*", "enable": True}) + model = _TransformerWithLMHead(n_layers=2, dim=16) + calib_data = [torch.randint(0, 32, (2, 8))] + + mtq.quantize(model, config, forward_loop=lambda m: [m(batch) for batch in calib_data]) + + assert model.lm_head.input_quantizer._amax is not None + with torch.no_grad(): + model(calib_data[0]) + + @pytest.mark.parametrize( "algorithm", ["gptq", "awq_lite", "smoothquant", "mse"],