P5: Fix mhc_rmsnorm_quantize_nvfp4 — add proper function definition

This commit is contained in:
2026-06-02 17:57:33 +00:00
parent 36fdbeb56d
commit c926c4a597

View File

@@ -395,6 +395,23 @@ def dequantize_nvfp4(x_fp4, x_sf, gsa, shape=None):
return mod.dequant_nvfp4(x_fp4, x_sf, gsa)
def mhc_rmsnorm_quantize_nvfp4(X_l, A_l, norm_weight, eps=1e-6, divisor=6.0 * 448.0):
"""Fused mHC pre_block + RMSNorm + NVFP4 quantize: 2 kernel launches total.
Replaces: bmm (1 launch) + rmsnorm (4+ launches) + quantize (2 launches)
Total unfused: 7+ launches per site × 122 sites = 854+ launches/token
Fused: 2 launches per site × 122 sites = 244 launches → 610 launches saved/token.
Args:
X_l: (M, n_hc, N) BF16 tensor. n_hc must be <= 4, N multiple of 16.
A_l: (M, n_hc) BF16 tensor. Softmax weights from mHC._dynamic_params.
norm_weight: (N,) FP32 RMSNorm weight.
eps: RMSNorm epsilon (default 1e-6).
divisor: gsa = amax / divisor. Default 6.0 * 448.0 = 2688.0.
Returns:
QuantizedActivation with x_fp4, x_sf, gsa, inv_rms
"""
from dsv4.kernels.cuda.loader import get_cuda_module
mod = get_cuda_module("fused_mhc_rmsnorm_quantize", ["fused_mhc_rmsnorm_quantize.cu"])
x_fp4, x_sf, gsa, inv_rms = mod.mhc_rmsnorm_quantize_nvfp4(X_l, A_l, norm_weight, eps, divisor)