[NVBug 6143871] Fix awq_lite uncalibrated branch leaving input_quantizer disabled (#1410)

## Summary

`awq_lite.setup()` disables `module.input_quantizer` at the start of
search. The calibrated branch re-enables it inside `postprocess()`, but
the uncalibrated branch (no cache-pass tokens, e.g. an MoE expert that
never gets routed) never did. Worse, for experts that had cache hits but
missed the search pass, the per-channel `_amax` left over from
`max_calibrate` during cache mode tripped `preprocess_linear_fusion`'s
`numel == 1` assertion and prevented `_export_quantized_weight` from
emitting a per-tensor `input_scale`.

Result: per-expert `input_scale` was missing in the exported HF
checkpoint, and TRT-LLM `CutlassFusedMoE` crashed on load with
`KeyError: '<idx>.w1.input_scale'` for any expert that did not see
enough calibration tokens (e.g. Qwen3-30B-A3B + `nvfp4_awq` from the bug
report).

## Fix

Mirror the calibrated `postprocess()` path in
`modelopt/torch/quantization/model_calib.py`: collapse any per-channel
`_amax` to scalar (axis=None) and re-enable the `input_quantizer`.

## Test plan

- [x] Added regression test
`test_awq_lite_uncalibrated_linear_keeps_input_quantizer_enabled` using
`NVFP4_AWQ_LITE_CFG` with a two-branch model where only one branch is
exercised; verifies the uncalibrated linear's `input_quantizer` remains
enabled after `mtq.quantize`.
- [x] End-to-end pipeline test on a tiny synthetic Qwen3-MoE (8 experts,
top-1 routing) confirms all 48 expert `input_scale` keys are present in
the exported state_dict (vs. multiple missing pre-fix).
- [x] Manual repro of the bug command (Qwen3-30B-A3B, `--qformat
nvfp4_awq`) confirmed the missing `input_scale` keys before the fix.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **Bug Fixes**
* Fixed AWQ-Lite quantization for uncalibrated modules so
export/preprocessing invariants are preserved even when
calibration/parameter updates are skipped.

* **Tests**
* Added regression tests: one verifies uncalibrated linear modules keep
their input quantizer enabled after quantization; another verifies
weight-only AWQ-Lite cases keep the input quantizer disabled as
expected.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
This commit is contained in:
Chenjie Luo
2026-05-08 04:44:14 +00:00
committed by GitHub
parent 6a3b6b8329
commit e2d29c869b
2 changed files with 93 additions and 0 deletions
@@ -1316,6 +1316,21 @@ def awq_lite(
dtype=w_dtype,
device=w_device,
)
# Mirror the calibrated postprocess path, gated on
# is_input_quantized so weight-only AWQ configs (where
# setup() never disabled input_quantizer) stay untouched.
# Collapse any per-channel _amax left over from cache_mode
# max_calibrate into a per-tensor scalar so
# preprocess_linear_fusion's numel==1 assertion passes, and
# re-enable the quantizer (awq_lite.setup disabled it).
if module.awq_lite.is_input_quantized:
if module.input_quantizer.amax is not None:
act_amax = module.input_quantizer.amax
module.input_quantizer._amax_for_smoothing = act_amax.cpu()
module.input_quantizer.reset_amax()
module.input_quantizer.axis = None
module.input_quantizer.amax = act_amax.amax()
module.input_quantizer.enable()
else:
with enable_weight_access_and_writeback(module, model, name_to_module):
postprocess(module, name)
@@ -312,6 +312,84 @@ def test_padded_awq():
model(torch.randn(2, 16, 16))
class _TwoBranchModel(nn.Module):
"""Two parallel linears; only the first is exercised by forward_loop."""
def __init__(self):
super().__init__()
self.calibrated = nn.Linear(16, 16, bias=False)
self.uncalibrated = nn.Linear(16, 16, bias=False)
def forward(self, x, branch="calibrated"):
if branch == "calibrated":
return self.calibrated(x)
return self.uncalibrated(x)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="NVFP4 dynamic block quant is CUDA-only")
def test_awq_lite_uncalibrated_linear_keeps_input_quantizer_enabled():
"""Regression test for NVBug 6143871.
awq_lite.setup() disables the input_quantizer at the start of search. The
calibrated branch re-enables it inside postprocess(); the uncalibrated
branch (no cache-pass tokens, e.g. an MoE expert that never gets routed)
must do the same — otherwise downstream export (set_expert_quantizer_amax
+ _export_quantized_weight) drops the input_scale buffer and inference
runtimes that read per-expert input_scale (e.g. TRT-LLM CutlassFusedMoE)
crash with KeyError on '<idx>.w1.input_scale'.
Also asserts the export-critical scalar amax invariant (axis=None,
numel==1) — preprocess_linear_fusion enforces it for fused-expert groups.
"""
torch.manual_seed(0)
model = _TwoBranchModel().cuda()
def _forward_loop(m):
for _ in range(2):
m(torch.randn(2, 16, 16, device="cuda"), branch="calibrated")
mtq.quantize(model, mtq.NVFP4_AWQ_LITE_CFG, _forward_loop)
assert model.calibrated.input_quantizer.is_enabled
assert model.uncalibrated.input_quantizer.is_enabled, (
"Uncalibrated linear's input_quantizer must remain enabled after "
"awq_lite postprocess so export emits input_scale (NVBug 6143871)."
)
uncal_q = model.uncalibrated.input_quantizer
# When amax exists (cache-hit but search-miss path), it must be the
# scalar form export expects — preprocess_linear_fusion asserts numel==1.
# When it's None (truly never routed), set_expert_quantizer_amax will
# populate it during export.
if uncal_q.amax is not None:
assert uncal_q.axis is None
assert uncal_q.amax.numel() == 1
def test_awq_lite_uncalibrated_weight_only_keeps_input_quantizer_disabled():
"""Weight-only AWQ companion to NVBug 6143871.
For weight-only AWQ configs (input_quantizer disabled), awq_lite.setup()
never touches the input_quantizer, so the postprocess uncalibrated branch
must NOT enable it — doing so turns on quantization the user's config had
explicitly opted out of.
"""
torch.manual_seed(0)
model = _TwoBranchModel()
def _forward_loop(m):
for _ in range(2):
m(torch.randn(2, 16, 16), branch="calibrated")
mtq.quantize(model, mtq.INT4_AWQ_CFG, _forward_loop)
assert not model.calibrated.input_quantizer.is_enabled
assert not model.uncalibrated.input_quantizer.is_enabled, (
"Weight-only AWQ must not flip on the input_quantizer for "
"uncalibrated layers — that would silently quantize activations "
"the user's config left in full precision."
)
def test_smoothquant_enable_disable():
torch.manual_seed(1234)
model = _SimpleMLP()