Skip to content

Commit 79408c6

Browse files
committed
fix(model): reject unaligned FP8 QKV scale export
Signed-off-by: Yu Yao <yaoyu.094@gmail.com>
1 parent 846216e commit 79408c6

2 files changed

Lines changed: 42 additions & 1 deletion

File tree

src/megatron/bridge/models/conversion/model_bridge.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,7 @@
5454
from megatron.bridge.models.conversion.mapping_registry import MegatronMappingRegistry
5555
from megatron.bridge.models.conversion.param_mapping import (
5656
MegatronParamMapping,
57+
QKVMapping,
5758
)
5859
from megatron.bridge.models.conversion.peft_bridge import (
5960
AdapterWeight,
@@ -2118,6 +2119,21 @@ def build_export_fp8_tasks(
21182119
fp8_scale_inv_attr,
21192120
)
21202121

2122+
for global_name, fp8_flag in global_fp8_flags.items():
2123+
block_size = self._fp8_scale_block_size(fp8_flag)
2124+
if block_size is None:
2125+
continue
2126+
mapping = mapping_registry.megatron_to_hf_lookup(self._get_lora_unwrapped_name(global_name))
2127+
if not isinstance(mapping, QKVMapping):
2128+
continue
2129+
head_size = model_config.kv_channels or (model_config.hidden_size // model_config.num_attention_heads)
2130+
if head_size % block_size != 0:
2131+
raise ValueError(
2132+
f"Cannot export blockwise FP8 QKV scales: QKV head size {head_size} "
2133+
f"is not divisible by FP8 block size {block_size}. Scale blocks that cross "
2134+
"QKV head boundaries cannot be split into separate HF projections."
2135+
)
2136+
21212137
# 2) Expand the global name list with `*.scale_inv` entries.
21222138
# This defines the final deterministic task ordering.
21232139
expanded_global_names: list[str] = []

tests/unit_tests/models/test_fp8_param_export.py

Lines changed: 26 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@
3030
WeightConversionTask,
3131
_HFNameSuffixMapping,
3232
)
33-
from megatron.bridge.models.conversion.param_mapping import split_qkv_weights
33+
from megatron.bridge.models.conversion.param_mapping import QKVMapping, split_qkv_weights
3434
from megatron.bridge.models.hf_pretrained.causal_lm import PreTrainedCausalLM
3535

3636

@@ -160,6 +160,31 @@ def megatron_to_hf(self, megatron_weights, megatron_module):
160160

161161

162162
class TestFp8ParamExport:
163+
def test_build_export_fp8_tasks_rejects_unaligned_qkv_scale_blocks(self, monkeypatch):
164+
bridge = DummyBridge()
165+
global_name = "decoder.layers.0.self_attention.linear_qkv.weight"
166+
mapping = QKVMapping(global_name, "hf.q.weight", "hf.k.weight", "hf.v.weight")
167+
_patch_export_task_context(
168+
monkeypatch,
169+
bridge,
170+
global_name,
171+
registry_factory=lambda: MegatronMappingRegistry(mapping),
172+
detect_fp8=lambda *_a, **_k: {global_name: 128},
173+
)
174+
model = SimpleNamespace(
175+
config=SimpleNamespace(
176+
num_attention_heads=64,
177+
num_query_groups=8,
178+
hidden_size=2880,
179+
kv_channels=64,
180+
share_embeddings_and_output_weights=False,
181+
),
182+
named_parameters=lambda: [],
183+
)
184+
185+
with pytest.raises(ValueError, match="QKV head size 64.*FP8 block size 128"):
186+
bridge.build_export_fp8_tasks(SimpleNamespace(state=SimpleNamespace(source=SimpleNamespace())), [model])
187+
163188
@pytest.mark.parametrize(
164189
"export_weight_dtype, expect_unquantized",
165190
[("fp8", True), ("bf16", False)],

0 commit comments

Comments
 (0)