[Issue 543] [Bug fix] Fix dynamic input quant for AWQ (#726)

## What does this PR do?

**Type of change:** Bug fix <!-- Use one of the following: Bug fix, new
feature, new example, new tests, documentation. -->

**Overview:** Dynamic input quantizers, e.g., MXFP4, are not restored
after AWQ. This PR fix the issue.

## Usage
<!-- You can potentially add a usage example below. -->

```python
# Add a code snippet demonstrating how to use this
```

## Testing
<!-- Mention how have you tested your change if applicable. -->

Tested with MXFP4, NVFP4, int4

## Before your PR is "*Ready for review*"
<!-- If you haven't finished some of the above items you can still open
`Draft` PR. -->

- **Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)**
and your commits are signed.
- **Is this change backward compatible?**: Yes/No <!--- If No, explain
why. -->
- **Did you write any new necessary tests?**: Yes/No
- **Did you add or update any necessary documentation?**: Yes/No
- **Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**:
Yes/No <!--- Only for new features, API changes, critical bug fixes or
bw breaking changes. -->

## Additional Information
<!-- E.g. related issue. -->

Signed-off-by: weimingc <17592131+meenchen@users.noreply.github.com>
This commit is contained in:
Wei-Ming Chen
2025-12-30 12:06:40 -08:00
committed by GitHub
parent 883c8731aa
commit b655321d87
+12 -8
View File
@@ -751,14 +751,18 @@ def awq_lite(
delattr(module.weight_quantizer, "_pre_quant_scale")
if hasattr(module.input_quantizer, "_pre_quant_scale"):
delattr(module.input_quantizer, "_pre_quant_scale")
if module.awq_lite.is_input_quantized and module.input_quantizer.amax is not None:
act_amax = module.input_quantizer.amax
# TODO: make this a buffer after we support only heterogeneous checkpointing for MCore
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()
if module.awq_lite.is_input_quantized:
if module.input_quantizer.amax is not None:
act_amax = module.input_quantizer.amax
# TODO: make this a buffer after we support only heterogeneous checkpointing for MCore
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()
# for dynamic quantization, there is no amax, so we just enable the quantizer
else:
module.input_quantizer.enable()
if module.awq_lite.is_enabled:
apply_pre_quant_scale_and_smooth(module, 1.0 / module.awq_lite.best_scale)