fix(megamoe): normalize fp4 weight preparation contract

Separate source and destination FP4 scale packing in requant_fp4_to_gran_k so group16-to-group32 conversion always recomputes UE8M0 runtime scales by default.

Make prepare_fp4_weights_for_mega_moe accept raw grouped FP4 weights and scales, then perform optional requantization, DeepGEMM scale layout transform, and MegaMoE UTCCP weight transform internally.

Update the MegaMoE synthetic benchmark so baseline grouped GEMM uses runtime-layout weights while fused MegaMoE uses transformed weights from the same raw source tensors.

Tested: PYTHONPYCACHEPREFIX=/private/tmp/deepgemm_pycache python3 -m py_compile deep_gemm/__init__.py deep_gemm/mega/__init__.py deep_gemm/utils/math.py tests/test_layout.py tests/test_mega_moe.py

Tested: git diff --check

Not-tested: CUDA build, SM100/B300 runtime, and GLM-5.2 accuracy validation are not available locally.
This commit is contained in:
LuminolT
2026-07-08 18:48:04 +08:00
parent 2c7543130b
commit 007c645f87
5 changed files with 86 additions and 27 deletions
+35 -11
View File
@@ -1,6 +1,6 @@
import torch
from typing import Tuple, Optional
from ..utils.math import align, requant_fp4_to_gran_k
from ..utils.math import align, requant_fp4_to_gran_k, unpack_ue8m0_from_int
# noinspection PyBroadException
try:
@@ -176,21 +176,45 @@ def transform_weights_for_mega_moe(
return l1_weights, l2_weights
def prepare_fp4_weights_for_mega_moe(
l1_weights: Tuple[torch.Tensor, torch.Tensor],
l2_weights: Tuple[torch.Tensor, torch.Tensor],
source_weight_gran_k: int = 32,
runtime_weight_gran_k: int = 32,
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
def _prepare_raw_fp4_weight_for_mega_moe(
weights: Tuple[torch.Tensor, torch.Tensor],
source_weight_gran_k: int,
runtime_weight_gran_k: int,
source_scale_packed_ue8m0: bool,
) -> Tuple[torch.Tensor, torch.Tensor]:
weight, weight_sf = weights
if source_weight_gran_k != runtime_weight_gran_k:
if source_weight_gran_k != 16 or runtime_weight_gran_k != 32:
raise RuntimeError(
f'Unsupported MegaMoE FP4 weight granularity conversion: '
f'{source_weight_gran_k} -> {runtime_weight_gran_k}')
l1_weights = requant_fp4_to_gran_k(
l1_weights[0], l1_weights[1], source_weight_gran_k, runtime_weight_gran_k)
l2_weights = requant_fp4_to_gran_k(
l2_weights[0], l2_weights[1], source_weight_gran_k, runtime_weight_gran_k)
weight, weight_sf = requant_fp4_to_gran_k(
weight, weight_sf,
source_weight_gran_k, runtime_weight_gran_k,
src_scale_packed_ue8m0=source_scale_packed_ue8m0)
source_scale_packed_ue8m0 = False
if source_scale_packed_ue8m0:
weight_sf = unpack_ue8m0_from_int(weight_sf)
num_groups, mn, packed_k = weight.shape
weight_sf = _C.transform_sf_into_required_layout(
weight_sf, mn, packed_k * 2, (1, runtime_weight_gran_k), num_groups)
return weight, weight_sf
def prepare_fp4_weights_for_mega_moe(
l1_weights: Tuple[torch.Tensor, torch.Tensor],
l2_weights: Tuple[torch.Tensor, torch.Tensor],
source_weight_gran_k: int = 32,
runtime_weight_gran_k: int = 32,
source_scale_packed_ue8m0: bool = False,
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
l1_weights = _prepare_raw_fp4_weight_for_mega_moe(
l1_weights, source_weight_gran_k, runtime_weight_gran_k,
source_scale_packed_ue8m0)
l2_weights = _prepare_raw_fp4_weight_for_mega_moe(
l2_weights, source_weight_gran_k, runtime_weight_gran_k,
source_scale_packed_ue8m0)
return transform_weights_for_mega_moe(
l1_weights, l2_weights, weight_gran_k=runtime_weight_gran_k)