Image-2.1-Calibrated-NVFP4 / source /dynamic_scale.py
ProCreations's picture
Release calibrated Image2.1 NVFP4 transformer with dynamic scaling and BF16 rank correction, native SM120 runtime, quality evidence and real-time demo
1961af5 verified
Raw History Blame Contribute Delete
1.29 kB
"""Dynamic tensor-wide NVFP4 range, with FP32 reductions and correction rescaling."""
import torch,triton
import triton.language as tl
@triton.jit
def _partials(X,P,T,COUNT:tl.constexpr,K:tl.constexpr,B:tl.constexpr):
i=tl.program_id(0)*B+tl.arange(0,B)
x=tl.load(X+i,i<COUNT,0).to(tl.float32)
p=tl.load(P+i%K).to(tl.float32)
tl.store(T+tl.program_id(0),tl.max(tl.abs(x*p),0))
@triton.jit
def _finish(T,OLDG,OLDA,G,A,R,N:tl.constexpr,B:tl.constexpr):
i=tl.arange(0,B);v=tl.load(T+i,i<N,0)
g=2688./tl.maximum(tl.max(v,0),1.e-12)
oldg=tl.load(OLDG);olda=tl.load(OLDA)
tl.store(G,g);tl.store(A,olda*oldg/g);tl.store(R,g/oldg)
@triton.jit
def _upscale(U,R,V,N:tl.constexpr,B:tl.constexpr):
i=tl.program_id(0)*B+tl.arange(0,B)
u=tl.load(U+i,i<N,0).to(tl.float32);r=tl.load(R)
tl.store(V+i,u*r,i<N)
def scale(x,pre,gx,alpha,up):
count=x.numel();blocks=triton.cdiv(count,16384)
partial=torch.empty(blocks,device=x.device,dtype=torch.float32)
g=torch.empty_like(gx);a=torch.empty_like(alpha);r=torch.empty_like(gx)
_partials[(blocks,)](x,pre,partial,count,x.shape[-1],16384,num_warps=8)
_finish[(1,)](partial,gx,alpha,g,a,r,blocks,triton.next_power_of_2(blocks),num_warps=8)
u=torch.empty_like(up)
_upscale[(triton.cdiv(up.numel(),1024),)](up,r,u,up.numel(),1024)
return g,a,u