122 lines
5.5 KiB
Python
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()
|