Skip to content

Commit 5363203

Browse files
sugunav14claude
andcommitted
Address second review round: symmetric base_layer lookup, loud miss, value oracle
- _restore_qtensor_wrappers: normalize the .base_layer suffix on both the saved q_tensor_state keys and the module names, so the lookup works whether the model was compressed before adapters were attached (quantize.py --compress) or after (QATTrainer._quantize_model). Warn when saved weights match no module at all, which is the silent failure this PR set out to fix. - export.py: document why skipping the restore is safe (QATTrainer snapshots the state at trainer init from the base state from_pretrained already restored). - LoRA-QAT example test: compare exported scales against a direct PTQ export rather than only asserting key presence. - Unit tests: cover the reverse key direction and the no-match warning. - CHANGELOG: add the 0.47 Bug Fixes entry. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Suguna Velury <178320438+sugunav14@users.noreply.github.com>
1 parent e8d77d7 commit 5363203

6 files changed

Lines changed: 71 additions & 20 deletions

File tree

‎CHANGELOG.rst‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,8 @@ Changelog
2121

2222
**Bug Fixes**
2323

24+
- Fix QLoRA export in ``examples/llm_qat/export.py`` failing with ``AssertionError: Model already has modelopt state!`` (NVBug 6542481). The QLoRA training output is an adapter-only checkpoint, so ``from_pretrained`` resolves the quantized base model from ``adapter_config.json`` and ``enable_huggingface_checkpointing`` already restores its ModelOpt state; the export then restored a second time. It now restores only when the loaded model is not already converted. Two further breakages on the same path are also fixed: ``_restore_qtensor_wrappers`` matched no modules because PEFT re-parents the quantized linear as ``<name>.base_layer`` while ``q_tensor_state`` is keyed by the name it was saved with (the packed NVFP4 weight then reached ``F.linear`` and raised a shape error), and ``postprocess_state_dict`` silently dropped every ``base_layer.*`` key missing from a hand-maintained rename map — losing the NVFP4 ``weight_scale_2`` global scale and any linear ``bias`` (Qwen2-style q/k/v biases), and leaving ``base_layer`` in the exported AWQ ``pre_quant_scale`` key. The rename is now a generic ``.base_layer.`` strip.
25+
2426
0.46 (2026-08-xx)
2527
^^^^^^^^^^^^^^^^^
2628

‎examples/llm_qat/export.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,9 @@ def get_model(
5050

5151
# Restore modelopt state for LoRA models. For QAT/QAD models from_pretrained call handles this.
5252
# For QLoRA the base checkpoint is quantized, so from_pretrained already restored the state.
53+
# Skipping the restore below is only safe because QATTrainer writes modelopt_state_train.pth at
54+
# trainer init from that same base state, so it carries no quantizer values from_pretrained did
55+
# not already load. Revisit this if the trainer ever snapshots state later in training.
5356
if hasattr(model, "peft_config") and not ModeloptStateManager.is_converted(model):
5457
modelopt_state = mto.load_modelopt_state(f"{ckpt_path}/modelopt_state_train.pth")
5558
restore_from_modelopt_state(model, modelopt_state)

‎modelopt/torch/opt/plugins/transformers.py‎

Lines changed: 26 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -110,18 +110,38 @@ def _restore_qtensor_wrappers(model, model_path):
110110
q_tensor_state = mode_config.get("metadata", {}).get("q_tensor_state", {})
111111
if not q_tensor_state:
112112
continue
113-
for name, module in model.named_modules():
114-
if not isinstance(module, RealQuantLinear) or isinstance(module.weight, QTensorWrapper):
115-
continue
116-
# PEFT renames the quantized linear to `<name>.base_layer`, but `q_tensor_state` is
117-
# keyed by the name it was saved with, so fall back to the stripped name.
118-
key = name if name in q_tensor_state else name.removesuffix(".base_layer")
113+
# PEFT nests the quantized linear as `<name>.base_layer`, and either side may carry that
114+
# suffix: the state is saved without it when the base model was compressed before adapters
115+
# were attached (`quantize.py --compress`) and with it when compressed after
116+
# (`QATTrainer._quantize_model`). Normalize both so the lookup works in either direction.
117+
# This assumes transformers injects adapters in place; moving the trainer to
118+
# `get_peft_model` would prefix module names with `base_model.model.` and break it.
119+
q_tensor_state = {k.removesuffix(".base_layer"): v for k, v in q_tensor_state.items()}
120+
121+
pending = [
122+
(name, module)
123+
for name, module in model.named_modules()
124+
if isinstance(module, RealQuantLinear) and not isinstance(module.weight, QTensorWrapper)
125+
]
126+
matched = 0
127+
for name, module in pending:
128+
key = name.removesuffix(".base_layer")
119129
if key not in q_tensor_state:
120130
continue
121131
module._parameters["weight"] = QTensorWrapper(
122132
qtensor=module.weight.data,
123133
metadata=q_tensor_state[key]["metadata"],
124134
)
135+
matched += 1
136+
137+
# A total miss means the module names were remapped by a wrapper we do not know about.
138+
# Warn here rather than letting it surface as an opaque shape error during dequantization.
139+
if pending and not matched:
140+
warnings.warn(
141+
f"Found {len(q_tensor_state)} compressed weight(s) in {modelopt_state_path} but "
142+
f"re-wrapped none of the {len(pending)} candidate module(s); their names may have "
143+
"been remapped. The model will likely fail when the packed weights are used."
144+
)
125145

126146

127147
def _new_from_pretrained(cls, /, pretrained_model_name_or_path, *args, **kwargs):

‎tests/examples/llm_qat/test_llm_qat.py‎

Lines changed: 13 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -190,17 +190,22 @@ def test_qwen3_lora_qat_nvfp4(tiny_qwen3_path, tmp_path):
190190
with open(base_model_dir / "hf_quant_config.json") as f:
191191
assert json.load(f)["quantization"]["quant_algo"] == "NVFP4"
192192

193-
# Scales are derived from the calibrated amaxes; missing ones mean the export silently
194-
# fell back to uncalibrated quantizers.
195193
base_weights = load_file(base_model_dir / "model.safetensors")
196-
weight_scales = [k for k in base_weights if k.endswith(".weight_scale")]
197-
assert weight_scales, "no NVFP4 weight scales in the exported base model"
198-
for key in weight_scales:
199-
prefix = key.removesuffix(".weight_scale")
200-
assert f"{prefix}.weight_scale_2" in base_weights
201-
assert f"{prefix}.input_scale" in base_weights
202194
assert not any("base_layer" in k or k.endswith("_amax") for k in base_weights)
203195

196+
# LoRA freezes the base model, so exporting the PTQ checkpoint directly is a trusted oracle
197+
# for every calibrated value. Comparing against it catches scales that survive as keys but
198+
# were silently reset to defaults, which key-presence assertions alone would miss.
199+
ptq_export_dir = tmp_path / "ptq_export"
200+
_run_export(str(ptq_output_dir), str(ptq_export_dir))
201+
reference = load_file(ptq_export_dir / "model.safetensors")
202+
203+
scales = [k for k in reference if k.endswith(("_scale", "_scale_2"))]
204+
assert scales, "no NVFP4 scales in the reference PTQ export"
205+
for key in scales:
206+
assert key in base_weights, f"{key} missing from the LoRA-QAT export"
207+
assert torch.equal(base_weights[key], reference[key]), f"{key} does not match PTQ export"
208+
204209

205210
@pytest.mark.parametrize("backend", [
206211
"fsdp2",

‎tests/gpu/torch/export/test_export.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -291,7 +291,8 @@ def test_postprocess_state_dict_qlora_strips_base_layer():
291291
"layer1.bias",
292292
"layer1.pre_quant_scale",
293293
}
294-
assert processed_state_dict["layer1.weight_scale_2"] == torch.tensor([0.5])
294+
assert torch.equal(processed_state_dict["layer1.weight_scale_2"], torch.tensor([0.5]))
295+
assert torch.equal(processed_state_dict["layer1.bias"], torch.arange(4.0))
295296

296297

297298
@pytest.mark.parametrize(

‎tests/unit/torch/opt/plugins/test_hf_patching.py‎

Lines changed: 25 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -69,14 +69,22 @@ def __init__(self, base_layer):
6969
self.base_layer = base_layer
7070

7171

72-
def _compressed_model_and_state_dir(tmp_path):
72+
def _compressed_model_and_state_dir(tmp_path, state_keyed_with_base_layer=False):
7373
model = nn.Sequential()
7474
model.fc = nn.Linear(64, 32)
7575
mtq.quantize(model, mtq.NVFP4_DEFAULT_CFG, lambda m: m(torch.randn(2, 64)))
7676
mtq.compress(model)
7777
assert isinstance(model.fc.weight, QTensorWrapper)
7878

79-
torch.save(mto.modelopt_state(model), tmp_path / "modelopt_state.pth")
79+
state = mto.modelopt_state(model)
80+
if state_keyed_with_base_layer:
81+
# Compressing after the adapters are attached (QATTrainer._quantize_model) saves the
82+
# keys with the peft suffix already in them.
83+
for _, mode_config in state["modelopt_state_dict"]:
84+
q_tensor_state = mode_config.get("metadata", {}).get("q_tensor_state", {})
85+
for key in list(q_tensor_state):
86+
q_tensor_state[f"{key}.base_layer"] = q_tensor_state.pop(key)
87+
torch.save(state, tmp_path / "modelopt_state.pth")
8088

8189
# transformers>=5 loads weights by assigning a plain Parameter holding the already-packed
8290
# data, which leaves the module without its QTensorWrapper.
@@ -88,9 +96,10 @@ def _compressed_model_and_state_dir(tmp_path):
8896

8997

9098
@pytest.mark.parametrize("wrap_in_lora", [False, True])
91-
def test_restore_qtensor_wrappers(tmp_path, wrap_in_lora):
92-
"""`q_tensor_state` is keyed by the pre-peft name, so `<name>.base_layer` must still match."""
93-
model = _compressed_model_and_state_dir(tmp_path)
99+
@pytest.mark.parametrize("state_keyed_with_base_layer", [False, True])
100+
def test_restore_qtensor_wrappers(tmp_path, wrap_in_lora, state_keyed_with_base_layer):
101+
"""Either side may carry the `.base_layer` suffix, so the lookup must work in both directions."""
102+
model = _compressed_model_and_state_dir(tmp_path, state_keyed_with_base_layer)
94103
if wrap_in_lora:
95104
model.fc = _LoraLike(model.fc)
96105

@@ -99,3 +108,14 @@ def test_restore_qtensor_wrappers(tmp_path, wrap_in_lora):
99108
linear = model.fc.base_layer if wrap_in_lora else model.fc
100109
assert isinstance(linear.weight, QTensorWrapper)
101110
assert linear.weight.metadata["shape"] == torch.Size([32, 64])
111+
112+
113+
def test_restore_qtensor_wrappers_warns_when_nothing_matches(tmp_path):
114+
"""A total miss must be loud -- it otherwise surfaces as an opaque shape error at dequant."""
115+
model = _compressed_model_and_state_dir(tmp_path)
116+
model.fc = _LoraLike(_LoraLike(model.fc)) # a nesting the lookup does not know about
117+
118+
with pytest.warns(UserWarning, match="re-wrapped none"):
119+
_restore_qtensor_wrappers(model, str(tmp_path))
120+
121+
assert not isinstance(model.fc.base_layer.base_layer.weight, QTensorWrapper)

0 commit comments

Comments
 (0)