Fix OOB read in mxfp4 MoE weight scales on Hopper (#36456)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
co-authored by
Mohammad Angkad
parent
e7e7894016
commit
2935bb8e79
@@ -213,6 +213,16 @@ if _is_hip:
|
|||||||
dynamic_mxfp4_quant = e8m0_shuffle = err
|
dynamic_mxfp4_quant = e8m0_shuffle = err
|
||||||
|
|
||||||
|
|
||||||
|
def _pad_hopper_mxfp4_scale(scale, k_size):
|
||||||
|
# triton_kernels' HOPPER_SCALE branch (matmul_details/_matmul.py) loads the w
|
||||||
|
# scales unmasked over cdiv(k_size, 128) tiles; drop when that load is masked.
|
||||||
|
mxfp4_block = 32
|
||||||
|
want = round_up(k_size, 128) // mxfp4_block
|
||||||
|
if scale.shape[-1] >= want:
|
||||||
|
return scale
|
||||||
|
return torch.nn.functional.pad(scale, (0, want - scale.shape[-1]), value=_UE8M0_ONE)
|
||||||
|
|
||||||
|
|
||||||
def _swizzle_mxfp4(quant_tensor, scale, num_warps):
|
def _swizzle_mxfp4(quant_tensor, scale, num_warps):
|
||||||
"""weight swizzle for mxfp4 moe, used for OAI mxfp4 kernel"""
|
"""weight swizzle for mxfp4 moe, used for OAI mxfp4 kernel"""
|
||||||
import triton_kernels.matmul_details.opt_flags as opt_flags
|
import triton_kernels.matmul_details.opt_flags as opt_flags
|
||||||
@@ -237,6 +247,8 @@ def _swizzle_mxfp4(quant_tensor, scale, num_warps):
|
|||||||
"split_k": 1,
|
"split_k": 1,
|
||||||
}
|
}
|
||||||
opt_flags.update_opt_flags_constraints(constraints)
|
opt_flags.update_opt_flags_constraints(constraints)
|
||||||
|
k_size = quant_tensor.shape[-1] * 2 # packed e2m1: 2 fp4 values per byte
|
||||||
|
scale = _pad_hopper_mxfp4_scale(scale=scale, k_size=k_size)
|
||||||
# transpose the tensor so that the quantization axis is on dim1
|
# transpose the tensor so that the quantization axis is on dim1
|
||||||
quant_tensor = quant_tensor.transpose(-2, -1)
|
quant_tensor = quant_tensor.transpose(-2, -1)
|
||||||
scale = scale.transpose(-2, -1)
|
scale = scale.transpose(-2, -1)
|
||||||
|
|||||||
Reference in New Issue
Block a user