feat: NVFP4 mega MoE kernel (scale_vec::4X, UE4M3 block scales)
- New CUDA kernel: sm100_fp8_nvfp4_mega_moe_impl - kGranK=16 (NVFP4 group_size=16, vs MXFP4's 32) - kind::mxf4nvf4.block_scale.scale_vec::4X PTX instruction - float_ue4m3_t scale factor type in instruction descriptor - SF layout: scale_vec::4X (4 TMEM sub-columns per UMMA atom) - UTCCP column stride: i*8 (vs MXFP4's i*4) for 4X layout - L1 epilogue: UE4M3 activation scales (float→cutlass::float_e4m3_t) - SF loading: kNumSFUint32 = kHidden/64 (4 UE4M3 per int32) - New PTX wrappers: SM100_MMA_MXF4NVF4_2x1SM_SS, SM100_MMA_MXF4NVF4_SS - Python API: - fp8_nvfp4_mega_moe() with recipe=(1,1,16) - transform_nvfp4_weights_for_mega_moe() for UE4M3→int32 UTCCP packing - _pack_nvfp4_sf_for_utccp() helper - C++ bindings: - mega_nvfp4.hpp with NVFP4-specific SymmBuffer (SF stride K/16) - JIT kernel header with kGranK=16 TMA descriptors - Registered in python_api.cpp NOTE: Both SFA and SFB must use UE4M3 (scale_format_ is 1-bit, shared). The L1 epilogue converts float→UE4M3 for activation scales.
This commit is contained in:
@@ -105,6 +105,57 @@ def transform_weights_for_mega_moe(
|
||||
return l1_weights, l2_weights
|
||||
|
||||
|
||||
def _pack_nvfp4_sf_for_utccp(sf: torch.Tensor) -> torch.Tensor:
|
||||
"""Pack NVFP4 UE4M3 block scales (float8_e4m3fn) into int32 UTCCP layout.
|
||||
|
||||
NVFP4 uses UE4M3 scales with group_size=16 (scale_vec::4X).
|
||||
The UTCCP layout packs 4 consecutive scale bytes into each int32,
|
||||
then applies the 4x32 transpose for TMA consumption.
|
||||
|
||||
Input: (num_experts, mn, K//16) float8_e4m3fn scales
|
||||
Output: (num_experts, mn, K//64) int32 packed UTCCP-transposed scales
|
||||
"""
|
||||
num_groups, mn, sf_k = sf.shape
|
||||
assert sf_k % 4 == 0, f"NVFP4 SF K dim must be divisible by 4, got {sf_k}"
|
||||
assert mn % 128 == 0, f"MN dim must be divisible by 128, got {mn}"
|
||||
|
||||
# View as uint8 and pack 4 consecutive bytes into int32
|
||||
sf_uint8 = sf.view(torch.uint8) # (num_groups, mn, sf_k)
|
||||
# Pack: every 4 uint8 → 1 int32
|
||||
packed = (sf_uint8[..., 0::4].to(torch.int32) |
|
||||
(sf_uint8[..., 1::4].to(torch.int32) << 8) |
|
||||
(sf_uint8[..., 2::4].to(torch.int32) << 16) |
|
||||
(sf_uint8[..., 3::4].to(torch.int32) << 24)) # (num_groups, mn, sf_k//4)
|
||||
|
||||
# Apply UTCCP 4x32 transpose (same as MXFP4 — the transpose is determined
|
||||
# by the 128-element alignment, not the scale vector size)
|
||||
packed_sf_k = sf_k // 4
|
||||
result = (packed.reshape(num_groups, -1, 4, 32, packed_sf_k)
|
||||
.transpose(2, 3)
|
||||
.reshape(num_groups, mn, packed_sf_k))
|
||||
return torch.empty_like(packed).copy_(result)
|
||||
|
||||
|
||||
def transform_nvfp4_weights_for_mega_moe(
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor]
|
||||
) -> Tuple[Tuple[torch.Tensor, torch.Tensor], Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""Transform NVFP4 expert weights for the mega_moe kernel.
|
||||
|
||||
NVFP4 weights come as (weight, scale) where:
|
||||
- weight: uint8 E2M1 packed, shape (num_experts, N, K//2)
|
||||
- scale: float8_e4m3fn UE4M3 block scales, shape (num_experts, N, K//16)
|
||||
|
||||
The kernel expects (weight, packed_sf) where packed_sf is int32 UTCCP layout.
|
||||
"""
|
||||
# L1: interleave gate/up, then pack + transpose SF for UTCCP
|
||||
l1_interleaved = _interleave_l1_weights(l1_weights)
|
||||
l1_weights = (l1_interleaved[0], _pack_nvfp4_sf_for_utccp(l1_interleaved[1]))
|
||||
# L2: only pack + transpose SF for UTCCP
|
||||
l2_weights = (l2_weights[0], _pack_nvfp4_sf_for_utccp(l2_weights[1]))
|
||||
return l1_weights, l2_weights
|
||||
|
||||
|
||||
def fp8_fp4_mega_moe(y: torch.Tensor,
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
@@ -126,3 +177,32 @@ def fp8_fp4_mega_moe(y: torch.Tensor,
|
||||
activation, activation_clamp,
|
||||
fast_math
|
||||
)
|
||||
|
||||
|
||||
def fp8_nvfp4_mega_moe(y: torch.Tensor,
|
||||
l1_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
l2_weights: Tuple[torch.Tensor, torch.Tensor],
|
||||
sym_buffer: SymmBuffer,
|
||||
cumulative_local_expert_recv_stats: Optional[torch.Tensor] = None,
|
||||
recipe: Tuple[int, int, int] = (1, 1, 16),
|
||||
activation: str = 'swiglu',
|
||||
activation_clamp: Optional[float] = None,
|
||||
fast_math: bool = True):
|
||||
"""NVFP4 mega MoE: uses kind::mxf4nvf4.block_scale.scale_vec::4X
|
||||
with UE4M3 block scales (group_size=16).
|
||||
|
||||
Weight format: (uint8 E2M1 packed, int32 packed UTCCP UE4M3 scales)
|
||||
Recipe: (1, 1, 16) — kGranK=16 for NVFP4 group_size=16.
|
||||
"""
|
||||
_C.fp8_nvfp4_mega_moe(
|
||||
y,
|
||||
l1_weights, l2_weights,
|
||||
cumulative_local_expert_recv_stats,
|
||||
sym_buffer.buffer,
|
||||
sym_buffer.handle.buffer_ptrs, sym_buffer.group.rank(),
|
||||
sym_buffer.num_max_tokens_per_rank,
|
||||
sym_buffer.num_experts, sym_buffer.num_topk,
|
||||
recipe,
|
||||
activation, activation_clamp,
|
||||
fast_math
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user