WaveCut's picture
Publish OrbitQuant 0.9.5 byte-pair decode and optional RMS activation fusion
b35ab1f verified
Raw History Blame Contribute Delete
3.2 kB
import json
from pathlib import Path
import pytest
import torch
from kernels import get_local_kernel
from orbitquant.kernels.native_packed_matmul import matmul_packed_w4a4_int8_with_native_kernel as original
kernel=get_local_kernel(Path(__file__).resolve().parents[1] / 'build', 'yue2_qkv_fused')
@pytest.mark.parametrize('rows,n,k',[(1,2048,2048),(2,6144,2048),(8,2048,6144),(1,129,128),(2,1024,2048),(1,129,130),(1,129,16)])
@pytest.mark.parametrize('dtype',[torch.bfloat16,torch.float16])
def test_exact_integer_gemv(rows,n,k,dtype):
torch.manual_seed(123)
x=torch.randint(0,256,(rows,k//2),device='cuda',dtype=torch.uint8)
w=torch.randint(0,256,(n*k//2,),device='cuda',dtype=torch.uint8)
xn=torch.rand(rows,device='cuda');wn=torch.rand(n,device='cuda',dtype=torch.bfloat16)
ac=torch.randint(-127,128,(16,),device='cuda',dtype=torch.int8);wc=ac.flip(0).contiguous()
bias=None
kw=dict(activation_scale=.03125,weight_scale=.0625,bias=bias,output_dtype=dtype)
if k % 128 == 0:
ref=original(x,w,xn,wn,ac,wc,out_features=n,in_features=k,**kw)
else:
def unpack(z):return torch.stack([z & 15,z >> 4],dim=-1).reshape(z.shape[0],-1).long().cpu()
aa=ac.cpu().long()[unpack(x)];ww=wc.cpu().long()[unpack(w.reshape(n,k//2))]
sums=(aa @ ww.T).to(device='cuda',dtype=torch.float32)
ref=(sums*(xn[:,None]*wn.float()[None,:]*(kw['activation_scale']*kw['weight_scale']))).to(dtype)
out=kernel.gemv(x,w,xn,wn,ac,wc,**kw)
torch.testing.assert_close(out,ref,rtol=0,atol=0)
stream=torch.cuda.Stream();stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(stream):
for _ in range(3):kernel.gemv(x,w,xn,wn,ac,wc,**kw)
torch.cuda.current_stream().wait_stream(stream)
graph=torch.cuda.CUDAGraph()
with torch.cuda.graph(graph): captured=kernel.gemv(x,w,xn,wn,ac,wc,**kw)
graph.replay();torch.cuda.synchronize()
torch.testing.assert_close(captured,ref,rtol=0,atol=0)
def test_rejects_bad_weight_shape():
x=torch.zeros(1,64,device='cuda',dtype=torch.uint8)
with pytest.raises(RuntimeError,match='packed weight size'):
kernel.gemv(x,x.flatten(),torch.ones(1,device='cuda'),torch.ones(2,device='cuda',dtype=torch.bfloat16),torch.zeros(16,device='cuda',dtype=torch.int8),torch.zeros(16,device='cuda',dtype=torch.int8),activation_scale=1,weight_scale=1)
@pytest.mark.parametrize('offset',[1,2,3])
def test_misaligned_contiguous_packed_inputs(offset):
rows,n,k=1,129,128
torch.manual_seed(10)
x=torch.randint(0,256,(rows*k//2+offset,),device='cuda',dtype=torch.uint8)[offset:].view(rows,k//2)
w=torch.randint(0,256,(n*k//2+offset,),device='cuda',dtype=torch.uint8)[offset:]
xn=torch.rand(rows,device='cuda');wn=torch.rand(n,device='cuda',dtype=torch.bfloat16)
ac=torch.randint(-127,128,(16,),device='cuda',dtype=torch.int8);wc=ac.flip(0).contiguous()
kw=dict(activation_scale=.03125,weight_scale=.0625,bias=None,output_dtype=torch.bfloat16)
ref=original(x.contiguous().clone(),w.contiguous().clone(),xn,wn,ac,wc,out_features=n,in_features=k,**kw)
out=kernel.gemv(x,w,xn,wn,ac,wc,**kw)
torch.testing.assert_close(out,ref,rtol=0,atol=0)