Fix fused router: use run_nvfp4_fused_router wrapper, correct CuTe tensor API

- kernel wrapper converts torch tensors to CuTe tensors with mark_layout_dynamic
- test uses the wrapper instead of calling kernel.run() directly
- mat_b/scale_b are now torch tensors (converted inside wrapper)
This commit is contained in:
2026-06-01 09:19:48 +00:00
parent 483e759d53
commit 7b3f6cb13c
2 changed files with 26 additions and 58 deletions

View File

@@ -1024,8 +1024,8 @@ class Nvfp4FusedRouterKernel:
# ================================================================
def run_nvfp4_fused_router(
hidden_states: torch.Tensor, # [M, hidden_size] BF16
mat_b, # CuTe tensor: gate weight (NVFP4, blockscaled layout)
scale_b, # CuTe tensor: gate weight scale factors (FP8 E4M3)
mat_b: torch.Tensor, # [K_packed, E_packed] NVFP4 gate weight (K-major, torch tensor)
scale_b: torch.Tensor, # [K_sf, E_sf] FP8 E4M3 weight scales (torch tensor)
gsa, # Activation global scale (scalar or 1-elem tensor)
gsb_val: float, # Weight global scale value
e_bias: torch.Tensor, # [E] FP32
@@ -1052,18 +1052,15 @@ def run_nvfp4_fused_router(
act_gsa = gsa if gsa is not None else _gsa
act_gsb = gsb_val
# Create CuTe tensors for activation + scales
# A tensor: [K_packed, M, L] where K_packed = K/2 (2 elements per byte for FP4)
K_packed = K // 2
mat_a = cutlass_torch.from_dlpack(act_nvfp4)
mat_a = mat_a.mark_layout_dynamic()
# SFA tensor: [K_sf, M, L]
scale_a = cutlass_torch.from_dlpack(act_sf)
scale_a = scale_a.mark_layout_dynamic()
def to_cute(t):
ct = cutlass_torch.from_dlpack(t)
return ct.mark_layout_dynamic(leading_dim=cutlass_torch.get_leading_dim(t))
# e_bias must be a CuTe tensor
e_bias_cute = cutlass_torch.from_dlpack(e_bias)
e_bias_cute = e_bias_cute.mark_layout_dynamic()
mat_a_cute = to_cute(act_nvfp4)
mat_b_cute = to_cute(mat_b)
scale_a_cute = to_cute(act_sf)
scale_b_cute = to_cute(scale_b)
e_bias_cute = to_cute(e_bias)
# Number of experts from e_bias
E = e_bias.shape[0]
@@ -1071,10 +1068,8 @@ def run_nvfp4_fused_router(
# Output buffers
out_weights = torch.zeros(M, top_k, dtype=torch.float32, device=device)
out_ids = torch.zeros(M, top_k, dtype=torch.int32, device=device)
out_w_cute = cutlass_torch.from_dlpack(out_weights)
out_w_cute = out_w_cute.mark_layout_dynamic()
out_id_cute = cutlass_torch.from_dlpack(out_ids)
out_id_cute = out_id_cute.mark_layout_dynamic()
out_w_cute = to_cute(out_weights)
out_id_cute = to_cute(out_ids)
# MMA tiler: (128, 128, 64) for decode
mma_tiler_mnk = (128, 128, 64)
@@ -1086,7 +1081,7 @@ def run_nvfp4_fused_router(
top_k=top_k,
)
kernel.run(
mat_a, mat_b, scale_a, scale_b,
mat_a_cute, mat_b_cute, scale_a_cute, scale_b_cute,
e_bias_cute, out_w_cute, out_id_cute,
M, E, K, routed_scaling_factor, top_k,
)

View File

@@ -92,50 +92,23 @@ def test_fused_router():
# ---- Fused kernel ----
print("\n[3] Running fused NVFP4 GEMM + router epilogue...")
from dsv4.kernels.router.nvfp4_fused_router_kernel import Nvfp4FusedRouterKernel
import cutlass.cute as cute
import cutlass.torch as cutlass_torch
# Quantize activation
act_gs = float(hidden_states.float().abs().max()) / (6.0 * 448.0)
act_nvfp4, act_sf = quantize_activation_nvfp4(hidden_states, act_gs)
# CuTe tensors for A (activation)
mat_a = cutlass_torch.from_dlpack(act_nvfp4)
mat_a = mat_a.mark_layout_dynamic()
scale_a = cutlass_torch.from_dlpack(act_sf)
scale_a = scale_a.mark_layout_dynamic()
# CuTe tensors for B (weight) — from gate_lin
mat_b = gate_lin._mat_b
scale_b = gate_lin._scale_b
# e_bias CuTe tensor
e_bias_cute = cutlass_torch.from_dlpack(e_bias)
e_bias_cute = e_bias_cute.mark_layout_dynamic()
# Output buffers
out_weights = torch.zeros(M, top_k, dtype=torch.float32, device=device)
out_ids = torch.zeros(M, top_k, dtype=torch.int32, device=device)
out_w_cute = cutlass_torch.from_dlpack(out_weights)
out_w_cute = out_w_cute.mark_layout_dynamic()
out_id_cute = cutlass_torch.from_dlpack(out_ids)
out_id_cute = out_id_cute.mark_layout_dynamic()
kernel = Nvfp4FusedRouterKernel(
sf_vec_size=sf_vec_size,
mma_tiler_mnk=(128, 128, 64),
cluster_shape_mnk=(1, 1, 1),
top_k=top_k,
)
from dsv4.kernels.router.nvfp4_fused_router_kernel import run_nvfp4_fused_router
try:
kernel.run(
mat_a, mat_b, scale_a, scale_b,
e_bias_cute, out_w_cute, out_id_cute,
M, E, K, routed_scaling_factor, top_k,
fused_weights, fused_ids = run_nvfp4_fused_router(
hidden_states=hidden_states,
mat_b=gate_lin._mat_b,
scale_b=gate_lin._scale_b,
gsa=gate_lin._gsa_buf,
gsb_val=float(gate_lin._gsb),
e_bias=e_bias,
routed_scaling_factor=routed_scaling_factor,
top_k=top_k,
sf_vec_size=sf_vec_size,
)
print(" Fused kernel compilation and execution succeeded!")
print(f" Fused topk_ids: {fused_ids[0].tolist()}")
print(f" Fused topk_weights: {fused_weights[0].tolist()}")
except Exception as ex:
print(f" FUSED KERNEL FAILED: {ex}")
import traceback