[NVBug: 6000530] Fix AWQ crash for uncalibrated MoE experts (#1142)

## Summary
- Fixes NVBugs 6000530: `AttributeError: 'float' object has no attribute
'pow'` when running AWQ lite with `moe_calib_experts_ratio < 1.0` on MoE
models (e.g. Qwen3-30B-A3B).
- **Root cause**: When `moe_calib_experts_ratio=0.5`, some MoE experts
receive zero tokens during the AWQ cache phase, leaving `act_scale` as a
Python float `0.0` instead of a tensor. This causes two failures:
1. **Search phase crash**: Uncalibrated experts crash in `get_scale()`
because `float.pow()` doesn't exist.
2. **Export crash**: Calibrated experts have `pre_quant_scale` but
uncalibrated ones don't, causing `torch.stack()` to fail on mixed
`None`/tensor values in `preprocess_linear_fusion()`.
- **Fix**: Handle uncalibrated experts (`num_cache_steps == 0`) in two
stages:
1. **Before search**: Disable AWQ search (`is_enabled = False`) to
prevent `get_scale()` crash on float `act_scale`.
2. **During postprocessing**: Max calibrate weights and apply a neutral
(all-ones) `pre_quant_scale` so export can stack scaling factors
consistently across all experts. The `pre_quant_scale` buffer must be
registered outside `enable_weight_access_and_writeback` because HF
accelerate's `post_forward` hook drops newly-registered submodule
buffers.

## Test plan
- [x] Reproduce with `Qwen/Qwen3-30B-A3B`, `--qformat int4_awq`,
`--moe_calib_experts_ratio 0.5` — verify no crash during calibration and
export

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

---------

Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Chenjie Luo
2026-03-31 13:55:21 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 74a8694a6e
commit ada1e26ba0
+37 -9
View File
@@ -1179,6 +1179,17 @@ def awq_lite(
module.parallel_state.data_parallel_group,
)
# Disable AWQ search for uncalibrated experts (num_cache_steps == 0) to
# prevent get_scale() crash on float act_scale. Max calibration and neutral
# pre_quant_scale are applied in the postprocessing loop below.
for name, module in model.named_modules():
if (
is_quantized_linear(module)
and hasattr(module, "awq_lite")
and module.awq_lite.num_cache_steps == 0
):
module.awq_lite.is_enabled = False
AWQLiteHelper.cache_mode = False
print_rank_0("awq_lite: Searching parameters...")
with torch.no_grad():
@@ -1212,16 +1223,33 @@ def awq_lite(
for name, module in model.named_modules():
if hasattr(module, "awq_lite"):
if module.awq_lite.num_cache_steps == 0:
module.awq_lite.is_enabled = False
elif module.awq_lite.num_search_steps == 0:
module.awq_lite.is_enabled = False
warnings.warn(
"awq_lite: Calling `forward_loop(model)` the second time did not forward data through the"
f" {name}. Please provide a valid `forward_loop` function that can be used to"
" forward data through the model many times."
# Uncalibrated expert: max calibrate weights and apply neutral
# (all-ones) pre_quant_scale for export consistency.
# NOTE: ones_scale must be registered OUTSIDE enable_weight_access_and_writeback
# because HF accelerate post_forward drops newly-registered submodule buffers.
with enable_weight_access_and_writeback(module, model, name_to_module):
max_calibrate(module, lambda module: module.weight_quantizer(module.weight))
w_shape, w_dtype, w_device = (
module.weight.shape[1],
module.weight.dtype,
module.weight.device,
)
module.input_quantizer._enable_pre_quant_scale = True
module.input_quantizer.pre_quant_scale = torch.ones(
w_shape,
dtype=w_dtype,
device=w_device,
)
with enable_weight_access_and_writeback(module, model, name_to_module):
postprocess(module, name)
else:
if module.awq_lite.num_search_steps == 0:
module.awq_lite.is_enabled = False
warnings.warn(
"awq_lite: Calling `forward_loop(model)` the second time did not forward"
f" data through the {name}. Please provide a valid `forward_loop` function"
" that can be used to forward data through the model many times."
)
with enable_weight_access_and_writeback(module, model, name_to_module):
postprocess(module, name)
module.awq_lite.cleanup()
if not debug: