P5: Fix mhc_rmsnorm_quantize_nvfp4 — add proper function definition
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user