diff --git a/tests/layertest.py b/tests/layertest.py index 37e3cfe3..53932a54 100644 --- a/tests/layertest.py +++ b/tests/layertest.py @@ -271,20 +271,15 @@ def main(): expert_offsets = compute_expert_offsets(tokens_per_expert, len(expert_indices)) - # ── BF16 L1 reference (gate+up only) ── + # ── BF16 L1 reference (slot-major, matching kernel output) ── print("\n Running BF16 L1 reference...") - ref_l1 = torch.zeros(num_tokens, 6144, dtype=torch.bfloat16, device=DEVICE) - for t in range(num_tokens): - for k in range(top_k): - e = expert_ids[t, k].item() - w = expert_weights[t, k].item() - if e not in nvfp4_experts_bf16: - continue - x = hidden_states[t] - gate = x @ nvfp4_experts_bf16[e]["gate_proj"].T # (3072,) - up = x @ nvfp4_experts_bf16[e]["up_proj"].T # (3072,) - ref_l1[t] += w * torch.cat([gate, up]) - + ref_l1_parts = [] + for e in expert_indices: + for t in expert_token_lists[e]: + gate = hidden_states[t] @ nvfp4_experts_bf16[e]["gate_proj"].T + up = hidden_states[t] @ nvfp4_experts_bf16[e]["up_proj"].T + ref_l1_parts.append(torch.cat([gate, up])) + ref_l1 = torch.cat(ref_l1_parts, dim=0) # (num_slots, 6144) print(f" BF16 L1 ref: amax={ref_l1.abs().max():.4f} mean={ref_l1.float().mean():.6f}") del nvfp4_experts_bf16 @@ -296,11 +291,14 @@ def main(): print(f" Kernel L1: amax={kernel_l1.abs().max():.4f} mean={kernel_l1.float().mean():.6f}") # ── Compare ── + ref_flat = ref_l1.flatten() + kernel_flat = kernel_l1.flatten() + cosine = torch.nn.functional.cosine_similarity( - kernel_l1.flatten().unsqueeze(0).float(), - ref_l1.flatten().unsqueeze(0).float(), + kernel_flat.unsqueeze(0).float(), + ref_flat.unsqueeze(0).float(), ).item() - mse = (kernel_l1.float() - ref_l1.float()).pow(2).mean().item() + mse = (kernel_flat.float() - ref_flat.float()).pow(2).mean().item() print(f"\n{'=' * 70}") print(f" RESULT: cosine={cosine:.6f} MSE={mse:.6e}")