matilda-jev-fp4 / fp4_kernels.py
yue-maincode's picture
Upload validated MATILDA JEV FP4 model and Decision Index scores
c69aaec verified
Raw History Blame Contribute Delete
1.39 kB
"""Portable Triton unpacking for E2M1 + E4M3 block16 scaled FP4 weights."""
import torch
import triton
import triton.language as tl
@triton.jit
def _unpack(W, S, G, O, SIZE:tl.constexpr, BLOCK:tl.constexpr):
i=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK)
packed=tl.load(W+i//2, i<SIZE, other=0).to(tl.int32)
code=tl.where(i%2==0,packed&15,packed>>4)
magnitude=code&7
value=tl.where(magnitude<4,magnitude.to(tl.float32)*.5,
tl.where(magnitude<6,magnitude.to(tl.float32)-2,(magnitude.to(tl.float32)-4)*2))
value=tl.where((code&8)!=0,-value,value)
scale_byte=tl.load(S+i//16,i<SIZE,other=0).to(tl.int32)
mantissa=scale_byte&7
exponent=(scale_byte>>3)&15
power=((exponent+120)<<23).to(tl.float32,bitcast=True)
scale=tl.where(exponent==0,mantissa.to(tl.float32)*.001953125,(1+mantissa.to(tl.float32)*.125)*power)
result=value*(scale*tl.load(G))
tl.store(O+i,result,i<SIZE)
def dequantize_weight(weight, scale, global_scale):
assert weight.is_cuda and weight.is_contiguous() and scale.is_contiguous()
n,k2=weight.shape
result=torch.empty((n,k2*2),device=weight.device,dtype=torch.bfloat16)
_unpack[(triton.cdiv(result.numel(),4096),)](weight,scale.view(torch.uint8),global_scale,result,result.numel(),4096,
num_warps=4,enable_fp_fusion=False)
return result