From b13c1057f578ae21f3e2d71b666c5fdd11679c3c Mon Sep 17 00:00:00 2001 From: biondizzle Date: Tue, 2 Jun 2026 08:43:40 +0000 Subject: [PATCH] test: verify GEMM shape with production weight format --- tests/unit/test_gemm_shape.py | 66 ++++++++++++++++------------------- 1 file changed, 31 insertions(+), 35 deletions(-) diff --git a/tests/unit/test_gemm_shape.py b/tests/unit/test_gemm_shape.py index e1be240c..54191728 100644 --- a/tests/unit/test_gemm_shape.py +++ b/tests/unit/test_gemm_shape.py @@ -1,57 +1,52 @@ #!/usr/bin/env python3 -"""Quick test: verify GEMM output shape for NVFP4.""" -import torch -import sys +"""Verify GEMM output shape — use production weight format.""" +import torch, sys sys.path.insert(0, '/root/dsv4-nvfp4-workspace/kernel') -from dsv4.ops.gemm_runner import ( - warmup_compilation, run_nvfp4_grouped_gemm, -) +from dsv4.ops.gemm_runner import warmup_compilation, run_nvfp4_grouped_gemm from dsv4.ops.quantize import quantize_to_nvfp4, quantize_activation_nvfp4 -from dsv4.ops.layouts import ( - make_b_k_major, interleave_l1_weights, - pad_and_swizzle_single, ceil_div as cutedsl_ceil_div, - assemble_scales_3d_side, -) +from dsv4.ops.layouts import (make_b_k_major, interleave_l1_weights, + pad_and_swizzle_single, ceil_div as cutedsl_ceil_div, assemble_scales_3d_side) device = "cuda:0" -K = 7168 -N = 6144 # gate+up combined -K_packed = K // 2 -N_packed = N // 2 +K = 7168; N = 6144 # gate+up +K_packed = K // 2; N_packed = N // 2 -# Create random data -x_bf16 = torch.randn(128, K, dtype=torch.bfloat16, device=device) * 0.1 -w_bf16 = torch.randn(1, K, N, dtype=torch.bfloat16, device=device) * 0.1 +# Create weight in PRODUCTION format: (N, K) BF16 → quantize → (N_packed, K_packed) float4 +torch.manual_seed(42) +w_bf16 = torch.randn(N, K, dtype=torch.bfloat16, device=device) * 0.1 +w_fp4, w_sf, w_gs = quantize_to_nvfp4(w_bf16) # (N_packed, K_packed) float4 +print(f"w_fp4 shape: {tuple(w_fp4.shape)} dtype={w_fp4.dtype}") -# Quantize -_, _, x_gs = quantize_to_nvfp4(x_bf16) -x_fp4, x_sf = quantize_activation_nvfp4(x_bf16, x_gs) -w_bf16_t = w_bf16.permute(0, 2, 1).contiguous() -w_fp4, w_sf, w_gs = quantize_to_nvfp4(w_bf16_t) +# Production path: (N_packed, K_packed) → (1, K_packed, N_packed) → interleave → make_b_k_major if w_fp4.dtype == torch.uint8: w_fp4 = w_fp4.view(torch.float4_e2m1fn_x2) -w_fp4_il = interleave_l1_weights(w_fp4) -mat_b = make_b_k_major(w_fp4_il) - -print(f"x_fp4 shape: {tuple(x_fp4.shape)} dtype={x_fp4.dtype}") +w_ekn = w_fp4.unsqueeze(0).permute(0, 2, 1).contiguous() # (1, K_packed, N_packed) +print(f"w_ekn shape (after permute): {tuple(w_ekn.shape)}") +w_ekn = interleave_l1_weights(w_ekn) +mat_b = make_b_k_major(w_ekn) print(f"mat_b shape: {tuple(mat_b.shape)} dtype={mat_b.dtype}") -print(f"N_packed from mat_b.shape[2]: {mat_b.shape[2]}") -# Run GEMM +# Activation +x_bf16 = torch.randn(128, K, dtype=torch.bfloat16, device=device) * 0.1 +_, _, x_gs = quantize_to_nvfp4(x_bf16) +x_fp4, x_sf = quantize_activation_nvfp4(x_bf16, x_gs) + +# Warmup warmup_compilation(1, K_packed, N_packed, device) -padded_offsets = torch.tensor([128], dtype=torch.int32, device=device) +# Scales +padded_offsets = torch.tensor([128], dtype=torch.int32, device=device) K_sf = cutedsl_ceil_div(K, 16) padded_cols = cutedsl_ceil_div(K_sf, 4) * 4 scale_a_buf = torch.zeros(128, padded_cols, dtype=torch.float16, device=device).to(torch.float8_e4m3fn) scale_a_buf[:128, :x_sf.shape[1]] = x_sf scale_a = pad_and_swizzle_single(scale_a_buf).reshape(128, padded_cols) -scale_b = assemble_scales_3d_side(w_sf) - +scale_b = assemble_scales_3d_side([w_sf]) gsa = torch.full((1,), x_gs, dtype=torch.float32, device=device) gsb = torch.full((1,), w_gs, dtype=torch.float32, device=device) +# Pad activation x_padded = torch.zeros(128, K_packed, dtype=torch.uint8, device=device).view(torch.float4_e2m1fn_x2) x_padded.view(torch.uint8)[:128] = x_fp4.view(torch.uint8) @@ -61,6 +56,7 @@ out = run_nvfp4_grouped_gemm( expert_offsets=padded_offsets, global_scale_a=gsa, global_scale_b=gsb, ) -print(f"\nGEMM output shape: {tuple(out.shape)} dtype={out.dtype}") -print(f"Expected: (128, {N}) BF16") -print(f"Got: (128, {out.shape[1]}) BF16") +print(f"\nGEMM output: shape={tuple(out.shape)} dtype={out.dtype}") +print(f"N (BF16) = {N}, N_packed = {N_packed}") +print(f"n_dim = mat_b.shape[2] = {mat_b.shape[2]}") +print(f"Output columns = {out.shape[1]}")