[ModelOpt][PP] Keep BF16 shared experts out of the NVFP4 fusion so TP1 pipeline stages can load (#40628)
This commit is contained in:
@@ -1497,6 +1497,15 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
|||||||
def get_min_capability(cls) -> int:
|
def get_min_capability(cls) -> int:
|
||||||
return 80
|
return 80
|
||||||
|
|
||||||
|
def can_fuse_shared_expert(self) -> bool:
|
||||||
|
# A shared-expert body kept BF16 via exclude_modules cannot share the packed
|
||||||
|
# FP4 FusedMoE buffers. The shared_expert_gate is a separate linear (kept
|
||||||
|
# BF16 by e.g. Qwen3-Next NVFP4 checkpoints) and must not veto fusion.
|
||||||
|
return not any(
|
||||||
|
"shared_expert" in name and "shared_expert_gate" not in name
|
||||||
|
for name in self.exclude_modules
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def common_group_size(cfg: dict) -> int:
|
def common_group_size(cfg: dict) -> int:
|
||||||
"""Return the unique group_size across the config; raise if missing/mismatched."""
|
"""Return the unique group_size across the config; raise if missing/mismatched."""
|
||||||
|
|||||||
@@ -132,6 +132,28 @@ class TestModelOptNvfp4(CustomTestCase):
|
|||||||
use_per_token_activation=True,
|
use_per_token_activation=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_shared_expert_fusion_requires_matching_fp4_precision(self):
|
||||||
|
quantized_shared = ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True,
|
||||||
|
group_size=16,
|
||||||
|
)
|
||||||
|
bf16_shared = ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True,
|
||||||
|
group_size=16,
|
||||||
|
exclude_modules=["model.layers.*.mlp.shared_experts*"],
|
||||||
|
)
|
||||||
|
|
||||||
|
gate_only_bf16 = ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True,
|
||||||
|
group_size=16,
|
||||||
|
exclude_modules=["model.layers.*.mlp.shared_expert_gate"],
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertTrue(quantized_shared.can_fuse_shared_expert())
|
||||||
|
self.assertFalse(bf16_shared.can_fuse_shared_expert())
|
||||||
|
# Only the gate is BF16 (Qwen3-Next NVFP4): the FP4 body still fuses.
|
||||||
|
self.assertTrue(gate_only_bf16.can_fuse_shared_expert())
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user