Files
deepseek-v4-quant/patches/test_nvfp4_mega_moe.py

122 lines
5.5 KiB
Python

"""Minimal test for fp8_nvfp4_mega_moe kernel with synthetic data.
Run inside the vllm container:
python3 /patches/test_nvfp4_mega_moe.py
"""
import torch
import torch.distributed as dist
import os
import sys
def test_nvfp4_mega_moe():
# Use dimensions that satisfy all alignment requirements:
# - hidden and intermediate_hidden must be multiples of 128 and 64
# - block_m will be at least 32 (SMEM alignment: 32 * 64 = 2048 >= 1024)
num_experts = 2
num_tokens = 32 # must be multiple of alignment
top_k = 2
hidden = 512 # multiple of 128 and 64
intermediate_hidden = 1024 # multiple of 128 and 64
device = "cuda"
torch.cuda.set_device(0)
# Single-rank process group for SymmBuffer
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29501")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("WORLD_SIZE", "1")
if not dist.is_initialized():
dist.init_process_group("nccl")
group = dist.new_group()
from deep_gemm.mega import (
fp8_nvfp4_mega_moe,
get_symm_buffer_for_nvfp4_mega_moe,
transform_nvfp4_weights_for_mega_moe,
)
# --- Weights: random NVFP4 ---
w13_weight = torch.randint(0, 256, (num_experts, 2 * intermediate_hidden, hidden // 2),
dtype=torch.uint8, device=device).view(torch.int8)
w13_weight_scale = torch.randn(num_experts, 2 * intermediate_hidden, hidden // 16,
device=device).abs().clamp(0.1, 10.0).to(torch.float8_e4m3fn)
w13_weight_scale_2 = torch.ones(num_experts, device=device) # global scale = 1 for simplicity
w2_weight = torch.randint(0, 256, (num_experts, hidden, intermediate_hidden // 2),
dtype=torch.uint8, device=device).view(torch.int8)
w2_weight_scale = torch.randn(num_experts, hidden, intermediate_hidden // 16,
device=device).abs().clamp(0.1, 10.0).to(torch.float8_e4m3fn)
w2_weight_scale_2 = torch.ones(num_experts, device=device)
print("Transforming weights...")
l1_weights, l2_weights = transform_nvfp4_weights_for_mega_moe(
(w13_weight, w13_weight_scale),
(w2_weight, w2_weight_scale),
l1_weight_scale_2=w13_weight_scale_2,
l2_weight_scale_2=w2_weight_scale_2,
)
for name, t in [("l1_w", l1_weights[0]), ("l1_w_sf", l1_weights[1]),
("l2_w", l2_weights[0]), ("l2_w_sf", l2_weights[1])]:
print(f" {name}: dtype={t.dtype} shape={tuple(t.shape)} strides={t.stride()} contig={t.is_contiguous()}")
# --- Symm buffer ---
print("Creating symm buffer...")
symm_buffer = get_symm_buffer_for_nvfp4_mega_moe(
group, num_experts, num_tokens, top_k, hidden, intermediate_hidden)
for name, t in [("x", symm_buffer.x), ("x_sf", symm_buffer.x_sf),
("l1_acts", symm_buffer.l1_acts), ("l1_acts_sf", symm_buffer.l1_acts_sf),
("l2_acts", symm_buffer.l2_acts), ("l2_acts_sf", symm_buffer.l2_acts_sf)]:
print(f" symm_{name}: dtype={t.dtype} shape={tuple(t.shape)} strides={t.stride()}")
# --- Stage inputs (BF16 hidden_states → FP4 packed + UE4M3 scales) ---
print("Staging inputs...")
hidden_states = torch.randn(num_tokens, hidden, dtype=torch.bfloat16, device=device) * 0.5
topk_weights = torch.softmax(torch.randn(num_tokens, top_k, device=device), dim=-1)
topk_ids = torch.randint(0, num_experts, (num_tokens, top_k), device=device)
# Import the staging kernel directly (can't import from vllm model without full init)
import triton
import triton.language as tl
# Just manually pack a few tokens into FP4 for testing
# Write actual FP4 data to the activation buffer (random but valid packed E2M1)
symm_buffer.x[:num_tokens].copy_(
torch.randint(0, 256, (num_tokens, hidden // 2), dtype=torch.uint8, device=device).view(torch.int8))
# Write valid UE4M3 scales (random but non-zero)
# x_sf shape is (tokens, hidden//64) as int32 — each int32 = 4 packed UE4M3 bytes
# Just fill with simple non-zero int32 values (the data doesn't need to be
# perfectly valid UE4M3 for a launch test, just non-garbage)
symm_buffer.x_sf[:num_tokens].fill_(0x3C3C3C3C) # repeating 0x3C = ~0.5 in E4M3
# Write topk data directly
for i in range(num_tokens):
for j in range(top_k):
symm_buffer.topk_idx[i, j] = topk_ids[i, j].item()
symm_buffer.topk_weights[i, j] = topk_weights[i, j].item()
torch.cuda.synchronize()
print("Buffer populated with random FP4 data")
# --- Run kernel ---
y = torch.zeros(num_tokens, hidden, dtype=torch.bfloat16, device=device)
print("Calling fp8_nvfp4_mega_moe...", flush=True)
import signal
timed_out = False
def handler(signum, frame):
nonlocal timed_out
timed_out = True
raise TimeoutError("Kernel timeout")
signal.signal(signal.SIGALRM, handler)
signal.alarm(15) # 15 second timeout
try:
fp8_nvfp4_mega_moe(y, l1_weights, l2_weights, symm_buffer)
torch.cuda.synchronize()
signal.alarm(0)
print(f"SUCCESS! y stats: min={y.min().item():.4f} max={y.max().item():.4f} mean={y.mean().item():.4f} nonzero={torch.count_nonzero(y).item()}")
except TimeoutError:
print("TIMEOUT: kernel did not complete in 15s (GPU hang?)")
except Exception as e:
signal.alarm(0)
print(f"FAILED: {e}")
raise
if __name__ == "__main__":
test_nvfp4_mega_moe()