fix: slot-major token layout for grouped GEMM
Tokens must be laid out as [expert0_tokens | expert1_tokens | ...] for the 2Dx3D grouped GEMM. Each expert gets its own contiguous block of tokens. Scale factors split by expert offsets.
This commit is contained in:
@@ -145,37 +145,23 @@ def moe_forward_bf16(hidden_states, experts, expert_ids, expert_weights):
|
||||
|
||||
# ── CuTeDSL NVFP4 Kernel MoE Forward ──────────────────────────────────
|
||||
|
||||
def moe_forward_nvfp4_l1_only(hidden_states, nvfp4_tensors, layer_idx, expert_ids, expert_weights):
|
||||
"""Run MoE forward pass using the CuTeDSL NVFP4 kernel via bridge."""
|
||||
num_tokens, hidden_size = hidden_states.shape
|
||||
top_k = expert_ids.shape[1]
|
||||
|
||||
# Map expert IDs to local indices
|
||||
unique_experts = sorted(set(expert_ids.flatten().tolist()))
|
||||
num_experts = len(unique_experts)
|
||||
expert_map = {e: i for i, e in enumerate(unique_experts)}
|
||||
|
||||
# ── Step 1: Quantize activation ──
|
||||
x_fp4, x_sf, x_igs = quantize_to_nvfp4(hidden_states)
|
||||
|
||||
# ── Step 2: Load and quantize weights from checkpoint ──
|
||||
# Checkpoint weight is (N, K//2) uint8, scale is (N, K//16) float8_e4m3fn
|
||||
# We need to dequantize to BF16 first, then re-quantize with our pipeline
|
||||
# (the checkpoint format is the same NVFP4, but we need to use our quantizer
|
||||
# for the bridge to produce correct tensor layouts)
|
||||
#
|
||||
# Actually, we can load the checkpoint weights directly as float4_e2m1fn_x2
|
||||
# and the scales as float8_e4m3fn. Just need to reshape.
|
||||
def moe_forward_nvfp4_l1_only(slot_hidden, nvfp4_tensors, layer_idx, expert_indices, tokens_per_expert):
|
||||
"""Run L1 (gate+up) GEMM using CuTeDSL.
|
||||
|
||||
slot_hidden is already laid out slot-major: [expert0_tokens | expert1_tokens | ...]
|
||||
"""
|
||||
num_slots, hidden_size = slot_hidden.shape
|
||||
num_experts = len(expert_indices)
|
||||
|
||||
# Quantize activation
|
||||
x_fp4, x_sf, x_igs = quantize_to_nvfp4(slot_hidden)
|
||||
|
||||
# Load and quantize weights
|
||||
w_fp4_list = []
|
||||
w_sf_list = []
|
||||
w_gs_list = []
|
||||
|
||||
for e in unique_experts:
|
||||
# L1: gate + up fused → (2*3072, 3584) packed
|
||||
# For now, dequantize checkpoint to BF16 then re-quantize
|
||||
# This ensures the FP4 values match our quantization convention
|
||||
|
||||
for e in expert_indices:
|
||||
gate_w_key = f"layers.{layer_idx}.mlp.experts.{e}.gate_proj.weight"
|
||||
gate_sf_key = f"layers.{layer_idx}.mlp.experts.{e}.gate_proj.weight_scale"
|
||||
gate_gs_key = f"layers.{layer_idx}.mlp.experts.{e}.gate_proj.weight_scale_2"
|
||||
@@ -193,52 +179,45 @@ def moe_forward_nvfp4_l1_only(hidden_states, nvfp4_tensors, layer_idx, expert_id
|
||||
nvfp4_tensors[up_sf_key].to(DEVICE),
|
||||
nvfp4_tensors[up_gs_key].item(),
|
||||
)
|
||||
|
||||
# Fuse gate + up: (6144, 7168) → quantize as (K=7168, N=6144)
|
||||
# Kernel expects B: (experts, K, N) with K=hidden, N=intermediate
|
||||
fused_l1 = torch.cat([gate_w_bf16, up_w_bf16], dim=0) # (6144, 7168)
|
||||
# B is (K, N) where K=hidden=7168, N=6144
|
||||
l1_w_bf16 = fused_l1.T # (7168, 6144) — K=7168 is dim 0
|
||||
|
||||
# Fuse gate + up, transpose to (K=hidden, N=6144)
|
||||
fused = torch.cat([gate_w_bf16, up_w_bf16], dim=0) # (6144, 7168)
|
||||
l1_w_bf16 = fused.T # (7168, 6144)
|
||||
l1_w_fp4, l1_w_sf, l1_w_gs = quantize_weight_to_nvfp4(l1_w_bf16)
|
||||
|
||||
|
||||
w_fp4_list.append(l1_w_fp4)
|
||||
w_sf_list.append(l1_w_sf)
|
||||
w_gs_list.append(l1_w_gs)
|
||||
|
||||
# Stack weights and convert to K-major
|
||||
mat_b = torch.stack(w_fp4_list) # (experts, K//2, N) N-major
|
||||
mat_b = make_b_k_major(mat_b) # (experts, K//2, N) K-major
|
||||
# Stack and convert to K-major
|
||||
mat_b = make_b_k_major(torch.stack(w_fp4_list))
|
||||
|
||||
# Assemble scale factors
|
||||
scale_a = assemble_scales_2d_side(
|
||||
[x_sf[e*top_k:(e+1)*top_k] for e in range(num_experts)]
|
||||
)
|
||||
# scale_a: per-expert activation scales, split by expert offsets
|
||||
x_sf_parts = []
|
||||
offset = 0
|
||||
for tpe in tokens_per_expert:
|
||||
x_sf_parts.append(x_sf[offset:offset+tpe])
|
||||
offset += tpe
|
||||
scale_a = assemble_scales_2d_side(x_sf_parts)
|
||||
scale_b = assemble_scales_3d_side(w_sf_list)
|
||||
|
||||
# Expert offsets
|
||||
tokens_per_expert = [top_k] * num_experts # simplified: each expert gets top_k tokens
|
||||
expert_offsets = compute_expert_offsets(tokens_per_expert, num_experts)
|
||||
|
||||
# Global scales
|
||||
global_scale_a = torch.tensor([x_igs] * num_experts, dtype=torch.float32, device=DEVICE)
|
||||
global_scale_b = torch.tensor(w_gs_list, dtype=torch.float32, device=DEVICE)
|
||||
|
||||
# Run the kernel
|
||||
# Run kernel
|
||||
out = run_nvfp4_grouped_gemm(
|
||||
mat_a=x_fp4,
|
||||
mat_b=mat_b,
|
||||
scale_a=scale_a,
|
||||
scale_b=scale_b,
|
||||
mat_a=x_fp4, mat_b=mat_b,
|
||||
scale_a=scale_a, scale_b=scale_b,
|
||||
expert_offsets=expert_offsets,
|
||||
global_scale_a=global_scale_a,
|
||||
global_scale_b=global_scale_b,
|
||||
global_scale_a=global_scale_a, global_scale_b=global_scale_b,
|
||||
)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
# ── Main ───────────────────────────────────────────────────────────────
|
||||
|
||||
def main():
|
||||
torch.manual_seed(42)
|
||||
expert_indices = [0, 1, 2]
|
||||
@@ -270,6 +249,28 @@ def main():
|
||||
expert_ids = torch.tensor([[0, 1]] * num_tokens, dtype=torch.int32, device=DEVICE)
|
||||
expert_weights = torch.tensor([[0.6, 0.4]] * num_tokens, dtype=torch.float32, device=DEVICE)
|
||||
|
||||
# ── Build slot-based layout for grouped GEMM ──
|
||||
# The kernel expects activation laid out as [expert_0_tokens | expert_1_tokens | ...]
|
||||
# Each token can appear in multiple experts (top-k routing)
|
||||
num_slots = num_tokens * top_k
|
||||
slot_expert = expert_ids.flatten() # (num_slots,)
|
||||
|
||||
# Build per-expert token lists
|
||||
expert_token_lists = {e: [] for e in expert_indices}
|
||||
for t in range(num_tokens):
|
||||
for k in range(top_k):
|
||||
e = expert_ids[t, k].item()
|
||||
expert_token_lists[e].append(t)
|
||||
|
||||
tokens_per_expert = [len(expert_token_lists[e]) for e in expert_indices]
|
||||
|
||||
# Build slot-major activation: concat tokens for each expert
|
||||
slot_hidden = torch.cat([
|
||||
hidden_states[expert_token_lists[e]] for e in expert_indices
|
||||
], dim=0) # (num_slots, hidden_size)
|
||||
|
||||
expert_offsets = compute_expert_offsets(tokens_per_expert, len(expert_indices))
|
||||
|
||||
# ── BF16 L1 reference (gate+up only) ──
|
||||
print("\n Running BF16 L1 reference...")
|
||||
ref_l1 = torch.zeros(num_tokens, 6144, dtype=torch.bfloat16, device=DEVICE)
|
||||
@@ -291,7 +292,7 @@ def main():
|
||||
|
||||
# ── CuTeDSL NVFP4 L1 kernel ──
|
||||
print("\n Running CuTeDSL NVFP4 L1 kernel (first run compiles, ~1-2 min)...")
|
||||
kernel_l1 = moe_forward_nvfp4_l1_only(hidden_states, nvfp4_tensors, LAYER_IDX, expert_ids, expert_weights)
|
||||
kernel_l1 = moe_forward_nvfp4_l1_only(slot_hidden, nvfp4_tensors, LAYER_IDX, expert_indices, tokens_per_expert)
|
||||
print(f" Kernel L1: amax={kernel_l1.abs().max():.4f} mean={kernel_l1.float().mean():.6f}")
|
||||
|
||||
# ── Compare ──
|
||||
|
||||
Reference in New Issue
Block a user