submission 103264
J · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 535 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-103264?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:a803c97a58c4bdd688c5c4c362a4fcf35c8de1a4d01129c80bcbf94918d40f0d
license declaredunknown
license concludedunknown
authorsJ
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
__device__ __forceinline__ float fp8e4m3_to_float(uint8_t v) {shared-memory
extern __shared__ uint8_t sm[];vector-width = uint4
__device__ __forceinline__ uint4 load_uint4_cs(const void* ptr) { uint4 r; asm volatile ("ld.global.cs.v4.u32 {%0,%1,%2,%3},[%4];" : "=r"(r.x),"=r"(r.y),"=r"(r.z),"=r"(r.w) : "l"(ptr)); return r; }Kernel source
submission.py535 lines
import torch
import cutlass
from typing import Tuple
input_t = Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
output_t = torch.Tensor
CUDA_SOURCE = r"""
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
__device__ const uint32_t fp4_packed_lut_const[] = {
0x0,0x3800,0x3c00,0x3e00,0x4000,0x4200,0x4400,0x4600,0x8000,0xb800,0xbc00,0xbe00,0xc000,0xc200,0xc400,0xc600,
0x38000000,0x38003800,0x38003c00,0x38003e00,0x38004000,0x38004200,0x38004400,0x38004600,0x38008000,0x3800b800,0x3800bc00,0x3800be00,0x3800c000,0x3800c200,0x3800c400,0x3800c600,
0x3c000000,0x3c003800,0x3c003c00,0x3c003e00,0x3c004000,0x3c004200,0x3c004400,0x3c004600,0x3c008000,0x3c00b800,0x3c00bc00,0x3c00be00,0x3c00c000,0x3c00c200,0x3c00c400,0x3c00c600,
0x3e000000,0x3e003800,0x3e003c00,0x3e003e00,0x3e004000,0x3e004200,0x3e004400,0x3e004600,0x3e008000,0x3e00b800,0x3e00bc00,0x3e00be00,0x3e00c000,0x3e00c200,0x3e00c400,0x3e00c600,
0x40000000,0x40003800,0x40003c00,0x40003e00,0x40004000,0x40004200,0x40004400,0x40004600,0x40008000,0x4000b800,0x4000bc00,0x4000be00,0x4000c000,0x4000c200,0x4000c400,0x4000c600,
0x42000000,0x42003800,0x42003c00,0x42003e00,0x42004000,0x42004200,0x42004400,0x42004600,0x42008000,0x4200b800,0x4200bc00,0x4200be00,0x4200c000,0x4200c200,0x4200c400,0x4200c600,
0x44000000,0x44003800,0x44003c00,0x44003e00,0x44004000,0x44004200,0x44004400,0x44004600,0x44008000,0x4400b800,0x4400bc00,0x4400be00,0x4400c000,0x4400c200,0x4400c400,0x4400c600,
0x46000000,0x46003800,0x46003c00,0x46003e00,0x46004000,0x46004200,0x46004400,0x46004600,0x46008000,0x4600b800,0x4600bc00,0x4600be00,0x4600c000,0x4600c200,0x4600c400,0x4600c600,
0x80000000,0x80003800,0x80003c00,0x80003e00,0x80004000,0x80004200,0x80004400,0x80004600,0x80008000,0x8000b800,0x8000bc00,0x8000be00,0x8000c000,0x8000c200,0x8000c400,0x8000c600,
0xb8000000,0xb8003800,0xb8003c00,0xb8003e00,0xb8004000,0xb8004200,0xb8004400,0xb8004600,0xb8008000,0xb800b800,0xb800bc00,0xb800be00,0xb800c000,0xb800c200,0xb800c400,0xb800c600,
0xbc000000,0xbc003800,0xbc003c00,0xbc003e00,0xbc004000,0xbc004200,0xbc004400,0xbc004600,0xbc008000,0xbc00b800,0xbc00bc00,0xbc00be00,0xbc00c000,0xbc00c200,0xbc00c400,0xbc00c600,
0xbe000000,0xbe003800,0xbe003c00,0xbe003e00,0xbe004000,0xbe004200,0xbe004400,0xbe004600,0xbe008000,0xbe00b800,0xbe00bc00,0xbe00be00,0xbe00c000,0xbe00c200,0xbe00c400,0xbe00c600,
0xc0000000,0xc0003800,0xc0003c00,0xc0003e00,0xc0004000,0xc0004200,0xc0004400,0xc0004600,0xc0008000,0xc000b800,0xc000bc00,0xc000be00,0xc000c000,0xc000c200,0xc000c400,0xc000c600,
0xc2000000,0xc2003800,0xc2003c00,0xc2003e00,0xc2004000,0xc2004200,0xc2004400,0xc2004600,0xc2008000,0xc200b800,0xc200bc00,0xc200be00,0xc200c000,0xc200c200,0xc200c400,0xc200c600,
0xc4000000,0xc4003800,0xc4003c00,0xc4003e00,0xc4004000,0xc4004200,0xc4004400,0xc4004600,0xc4008000,0x4400b800,0xc400bc00,0xc400be00,0xc400c000,0xc400c200,0xc400c400,0xc400c600,
0xc6000000,0xc6003800,0xc6003c00,0xc6003e00,0xc6004000,0xc6004200,0xc6004400,0xc6004600,0xc6008000,0xc600b800,0xc600bc00,0xc600be00,0xc600c000,0xc600c200,0xc600c400,0xc600c600,
};
// Precomputed FP8 E4M3 to float lookup table - eliminates runtime computation
__device__ const float fp8_lut_const[256] = {
// e=0 subnormal (v=0-7)
0.0f, 0.001953125f, 0.00390625f, 0.005859375f, 0.0078125f, 0.009765625f, 0.01171875f, 0.013671875f,
// e=1 (v=8-15)
0.015625f, 0.017578125f, 0.01953125f, 0.021484375f, 0.0234375f, 0.025390625f, 0.02734375f, 0.029296875f,
// e=2 (v=16-23)
0.03125f, 0.03515625f, 0.0390625f, 0.04296875f, 0.046875f, 0.05078125f, 0.0546875f, 0.05859375f,
// e=3 (v=24-31)
0.0625f, 0.0703125f, 0.078125f, 0.0859375f, 0.09375f, 0.1015625f, 0.109375f, 0.1171875f,
// e=4 (v=32-39)
0.125f, 0.140625f, 0.15625f, 0.171875f, 0.1875f, 0.203125f, 0.21875f, 0.234375f,
// e=5 (v=40-47)
0.25f, 0.28125f, 0.3125f, 0.34375f, 0.375f, 0.40625f, 0.4375f, 0.46875f,
// e=6 (v=48-55)
0.5f, 0.5625f, 0.625f, 0.6875f, 0.75f, 0.8125f, 0.875f, 0.9375f,
// e=7 (v=56-63)
1.0f, 1.125f, 1.25f, 1.375f, 1.5f, 1.625f, 1.75f, 1.875f,
// e=8 (v=64-71)
2.0f, 2.25f, 2.5f, 2.75f, 3.0f, 3.25f, 3.5f, 3.75f,
// e=9 (v=72-79)
4.0f, 4.5f, 5.0f, 5.5f, 6.0f, 6.5f, 7.0f, 7.5f,
// e=10 (v=80-87)
8.0f, 9.0f, 10.0f, 11.0f, 12.0f, 13.0f, 14.0f, 15.0f,
// e=11 (v=88-95)
16.0f, 18.0f, 20.0f, 22.0f, 24.0f, 26.0f, 28.0f, 30.0f,
// e=12 (v=96-103)
32.0f, 36.0f, 40.0f, 44.0f, 48.0f, 52.0f, 56.0f, 60.0f,
// e=13 (v=104-111)
64.0f, 72.0f, 80.0f, 88.0f, 96.0f, 104.0f, 112.0f, 120.0f,
// e=14 (v=112-119)
128.0f, 144.0f, 160.0f, 176.0f, 192.0f, 208.0f, 224.0f, 240.0f,
// e=15 saturate (v=120-127)
448.0f, 448.0f, 448.0f, 448.0f, 448.0f, 448.0f, 448.0f, 448.0f,
// Negative values (v=128-255): same pattern negated
0.0f, -0.001953125f, -0.00390625f, -0.005859375f, -0.0078125f, -0.009765625f, -0.01171875f, -0.013671875f,
-0.015625f, -0.017578125f, -0.01953125f, -0.021484375f, -0.0234375f, -0.025390625f, -0.02734375f, -0.029296875f,
-0.03125f, -0.03515625f, -0.0390625f, -0.04296875f, -0.046875f, -0.05078125f, -0.0546875f, -0.05859375f,
-0.0625f, -0.0703125f, -0.078125f, -0.0859375f, -0.09375f, -0.1015625f, -0.109375f, -0.1171875f,
-0.125f, -0.140625f, -0.15625f, -0.171875f, -0.1875f, -0.203125f, -0.21875f, -0.234375f,
-0.25f, -0.28125f, -0.3125f, -0.34375f, -0.375f, -0.40625f, -0.4375f, -0.46875f,
-0.5f, -0.5625f, -0.625f, -0.6875f, -0.75f, -0.8125f, -0.875f, -0.9375f,
-1.0f, -1.125f, -1.25f, -1.375f, -1.5f, -1.625f, -1.75f, -1.875f,
-2.0f, -2.25f, -2.5f, -2.75f, -3.0f, -3.25f, -3.5f, -3.75f,
-4.0f, -4.5f, -5.0f, -5.5f, -6.0f, -6.5f, -7.0f, -7.5f,
-8.0f, -9.0f, -10.0f, -11.0f, -12.0f, -13.0f, -14.0f, -15.0f,
-16.0f, -18.0f, -20.0f, -22.0f, -24.0f, -26.0f, -28.0f, -30.0f,
-32.0f, -36.0f, -40.0f, -44.0f, -48.0f, -52.0f, -56.0f, -60.0f,
-64.0f, -72.0f, -80.0f, -88.0f, -96.0f, -104.0f, -112.0f, -120.0f,
-128.0f, -144.0f, -160.0f, -176.0f, -192.0f, -208.0f, -224.0f, -240.0f,
-448.0f, -448.0f, -448.0f, -448.0f, -448.0f, -448.0f, -448.0f, -448.0f
};
__device__ __forceinline__ float fp8e4m3_to_float(uint8_t v) {
int s = (v >> 7) & 1; int e = (v >> 3) & 0xF; int m = v & 7;
if (e == 0) { if (m == 0) return 0.0f; return (s ? -1.0f : 1.0f) * (m / 8.0f) * ldexpf(1.0f, -6); }
if (e == 15) return s ? -448.0f : 448.0f;
float r = ldexpf(1.0f + (m / 8.0f), e - 7); return s ? -r : r;
}
__device__ __forceinline__ uint4 load_uint4_cs(const void* ptr) { uint4 r; asm volatile ("ld.global.cs.v4.u32 {%0,%1,%2,%3},[%4];" : "=r"(r.x),"=r"(r.y),"=r"(r.z),"=r"(r.w) : "l"(ptr)); return r; }
// Prefetch to L2 cache
__device__ __forceinline__ void prefetch_l2(const void* ptr) { asm volatile ("prefetch.global.L2 [%0];" :: "l"(ptr)); }
__device__ __forceinline__ uint4 load_uint4_ca(const void* ptr) { uint4 r; asm volatile ("ld.global.ca.v4.u32 {%0,%1,%2,%3},[%4];" : "=r"(r.x),"=r"(r.y),"=r"(r.z),"=r"(r.w) : "l"(ptr)); return r; }
// Cache global (L2 only, bypass L1) - good for large working sets on B200
__device__ __forceinline__ uint4 load_uint4_cg(const void* ptr) { uint4 r; asm volatile ("ld.global.cg.v4.u32 {%0,%1,%2,%3},[%4];" : "=r"(r.x),"=r"(r.y),"=r"(r.z),"=r"(r.w) : "l"(ptr)); return r; }
// Non-coherent load via texture cache - potentially better for streaming
__device__ __forceinline__ uint4 load_uint4_nc(const void* ptr) { uint4 r; asm volatile ("ld.global.nc.v4.u32 {%0,%1,%2,%3},[%4];" : "=r"(r.x),"=r"(r.y),"=r"(r.z),"=r"(r.w) : "l"(ptr)); return r; }
// Using __ldg for read-only cache
__device__ __forceinline__ uint4 load_uint4_ldg(const uint4* ptr) { return __ldg(ptr); }
// =============================================================================
// ROUTE-SPECIFIC FP4 CONVERSION IMPLEMENTATIONS
// Each K route can have its own optimized proc_ptx implementation
// =============================================================================
// Original FP4 conversion using uint64 - works best for k7168 and k16384
__device__ __forceinline__ half2 cvt_fp4_to_half2_orig(uint32_t x) {
uint32_t res; uint64_t x64=(uint64_t)x;
asm volatile("{.reg .b32 r32; cvt.u32.u64 r32,%1; .reg .b8 b0; mov.b32 {b0,_,_,_},r32; cvt.rn.f16x2.e2m1x2 %0,b0;}":"=r"(res):"l"(x64));
return *reinterpret_cast<half2*>(&res);
}
// Original proc_ptx for k7168 and k16384 routes
__device__ __forceinline__ void proc_ptx_orig(half2& acc, uint32_t a, uint32_t b) {
acc=__hfma2(cvt_fp4_to_half2_orig(a&0xFF),cvt_fp4_to_half2_orig(b&0xFF),acc);
acc=__hfma2(cvt_fp4_to_half2_orig((a>>8)&0xFF),cvt_fp4_to_half2_orig((b>>8)&0xFF),acc);
acc=__hfma2(cvt_fp4_to_half2_orig((a>>16)&0xFF),cvt_fp4_to_half2_orig((b>>16)&0xFF),acc);
acc=__hfma2(cvt_fp4_to_half2_orig(a>>24),cvt_fp4_to_half2_orig(b>>24),acc);
}
// Optimized proc_ptx for k2048 route - uses mov.b32 byte extraction
__device__ __forceinline__ void proc_ptx_k2048(half2& acc, uint32_t a, uint32_t b) {
uint32_t ra0, rb0, ra1, rb1, ra2, rb2, ra3, rb3;
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(a), "r"(b));
acc = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), acc);
}
// =============================================================================
// ROUTE-SPECIFIC blk_ptx IMPLEMENTATIONS
// =============================================================================
// blk_ptx for k2048 route
__device__ __forceinline__ void blk_ptx_k2048(float ac[4], uint4 va[4], uint4 vb, uint8_t sfa[4][2], uint8_t sfb[2], const float* f8) {
float f=f8[sfb[0]]; half2 s[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][0]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_k2048(s[k],va[k].x,vb.x); proc_ptx_k2048(s[k],va[k].y,vb.y); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
f=f8[sfb[1]];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][1]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_k2048(s[k],va[k].z,vb.z); proc_ptx_k2048(s[k],va[k].w,vb.w); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
}
// blk_ptx for k7168 and k16384 routes (uses original proc_ptx)
__device__ __forceinline__ void blk_ptx(float ac[4], uint4 va[4], uint4 vb, uint8_t sfa[4][2], uint8_t sfb[2], const float* f8) {
float f=f8[sfb[0]]; half2 s[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][0]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(s[k],va[k].x,vb.x); proc_ptx_orig(s[k],va[k].y,vb.y); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
f=f8[sfb[1]];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][1]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(s[k],va[k].z,vb.z); proc_ptx_orig(s[k],va[k].w,vb.w); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
}
// Optimized blk_ptx with precomputed float scales (unused currently, kept for future use)
__device__ __forceinline__ void blk_ptx_fast(float ac[4], uint4 va[4], uint4 vb, float sc0[4], float sc1[4]) {
half2 s[4];
#pragma unroll
for(int k=0;k<4;++k) s[k]=__float2half2_rn(0.0f);
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(s[k],va[k].x,vb.x); proc_ptx_orig(s[k],va[k].y,vb.y); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc0[k];
#pragma unroll
for(int k=0;k<4;++k) s[k]=__float2half2_rn(0.0f);
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(s[k],va[k].z,vb.z); proc_ptx_orig(s[k],va[k].w,vb.w); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc1[k];
}
__device__ __forceinline__ void process_packed_4bytes(half2& acc, uint32_t a, uint32_t b, const uint32_t* l) {
uint32_t ra, rb;
ra=l[a&0xFF]; rb=l[b&0xFF]; acc=__hfma2(*(half2*)&ra,*(half2*)&rb,acc);
ra=l[(a>>8)&0xFF]; rb=l[(b>>8)&0xFF]; acc=__hfma2(*(half2*)&ra,*(half2*)&rb,acc);
ra=l[(a>>16)&0xFF]; rb=l[(b>>16)&0xFF]; acc=__hfma2(*(half2*)&ra,*(half2*)&rb,acc);
ra=l[a>>24]; rb=l[b>>24]; acc=__hfma2(*(half2*)&ra,*(half2*)&rb,acc);
}
__device__ __forceinline__ void blk_lut(float ac[4], uint4 va[4], uint4 vb, uint8_t sfa[4][2], uint8_t sfb[2], const uint32_t* l, const float* f8) {
float f=f8[sfb[0]]; half2 s[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][0]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { process_packed_4bytes(s[k],va[k].x,vb.x,l); process_packed_4bytes(s[k],va[k].y,vb.y,l); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
f=f8[sfb[1]];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[sfa[k][1]]*f; }
#pragma unroll
for(int k=0;k<4;++k) { process_packed_4bytes(s[k],va[k].z,vb.z,l); process_packed_4bytes(s[k],va[k].w,vb.w,l); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
}
// Optimized proc_ptx for k16384 route - uses mov.b32 byte extraction like k2048
__device__ __forceinline__ void proc_ptx_k16384(half2& acc, uint32_t a, uint32_t b) {
uint32_t ra0, rb0, ra1, rb1, ra2, rb2, ra3, rb3;
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(a), "r"(b));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(a), "r"(b));
acc = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), acc);
acc = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), acc);
}
// k2048: 256 threads, 32 rows per block, cache B in shared memory
extern "C" __global__ __launch_bounds__(256) void k2048(const uint8_t* __restrict__ A, const uint8_t* __restrict__ B, const uint8_t* __restrict__ SFA, const uint8_t* __restrict__ SFB, half* __restrict__ C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2) {
int tid=threadIdx.x;
extern __shared__ uint8_t sm[];
float* f8=(float*)sm;
uint8_t* s_B=(uint8_t*)(f8+256);
uint8_t* s_SFB=s_B+1024;
// Compute LUT in shared memory
f8[tid] = fp8e4m3_to_float((uint8_t)tid);
// Load B to shared memory (1024 bytes)
const uint8_t* B_src=B+blockIdx.y*Bs2;
*reinterpret_cast<uint32_t*>(s_B+tid*4)=*reinterpret_cast<const uint32_t*>(B_src+tid*4);
// Load SFB to shared memory (128 bytes)
const uint8_t* SFB_src=SFB+blockIdx.y*SFBs2;
if(tid<32) *reinterpret_cast<uint32_t*>(s_SFB+tid*4)=*reinterpret_cast<const uint32_t*>(SFB_src+tid*4);
__syncthreads();
int r0=blockIdx.x*32+(tid>>5)*4; if(r0>=M) return;
const uint8_t* Ar[4]; const uint8_t* Sr[4];
#pragma unroll
for(int k=0;k<4;++k) { int r=r0+k; Ar[k]=A+(r<M?r:0)*As0+blockIdx.y*As2; Sr[k]=SFA+(r<M?r:0)*SFAs0+blockIdx.y*SFAs2; }
float ac[4]={0,0,0,0}; int cur=(tid&31)*16;
// 2 iterations for K=2048 (1024 bytes / 512 bytes per iteration)
#pragma unroll
for(int iter=0; iter<2; ++iter) {
int off = cur + iter*512;
uint4 b=*reinterpret_cast<const uint4*>(s_B+off); uint4 va[4];
#pragma unroll
for(int k=0;k<4;++k) va[k]=load_uint4_cs(Ar[k]+off);
uint8_t sfa[4][2], sfb[2]; int si=off>>3;
#pragma unroll
for(int k=0;k<4;++k) { sfa[k][0]=Sr[k][si*SFAs1]; sfa[k][1]=Sr[k][(si+1)*SFAs1]; }
sfb[0]=s_SFB[si]; sfb[1]=s_SFB[si+1];
blk_ptx_k2048(ac,va,b,sfa,sfb,f8);
}
#pragma unroll
for(int o=16;o>0;o/=2) for(int k=0;k<4;++k) ac[k]+=__shfl_down_sync(0xffffffff,ac[k],o);
if((tid&31)==0) for(int k=0;k<4;++k) { int r=r0+k; if(r<M) C[r*Cs0+blockIdx.y*Cs2]=__float2half(ac[k]); }
}
// k7168: 128 threads, 4 rows per warp (16 rows/block) - optimal config
extern "C" __global__ __launch_bounds__(128) void k7168(const uint8_t* __restrict__ A, const uint8_t* __restrict__ B, const uint8_t* __restrict__ SFA, const uint8_t* __restrict__ SFB, half* __restrict__ C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2) {
int tid=threadIdx.x;
extern __shared__ uint8_t sm[];
float* f8=(float*)sm;
uint8_t* s_B=(uint8_t*)(f8+256);
float* s_SFB_f=(float*)(s_B+3584);
f8[tid] = fp8_lut_const[tid];
f8[tid+128] = fp8_lut_const[tid+128];
const uint8_t* B_src=B+blockIdx.y*Bs2;
*reinterpret_cast<uint4*>(s_B+tid*16)=*reinterpret_cast<const uint4*>(B_src+tid*16);
if(tid<96) *reinterpret_cast<uint4*>(s_B+(tid+128)*16)=*reinterpret_cast<const uint4*>(B_src+(tid+128)*16);
const uint8_t* SFB_src=SFB+blockIdx.y*SFBs2;
s_SFB_f[tid] = fp8_lut_const[SFB_src[tid]];
s_SFB_f[tid+128] = fp8_lut_const[SFB_src[tid+128]];
s_SFB_f[tid+256] = fp8_lut_const[SFB_src[tid+256]];
if(tid<64) s_SFB_f[tid+384] = fp8_lut_const[SFB_src[tid+384]];
__syncthreads();
int r0=blockIdx.x*16+(tid>>5)*4; if(r0>=M) return;
const uint8_t* Ar[4]; const uint8_t* Sr[4];
#pragma unroll
for(int k=0;k<4;++k) { int r=r0+k; Ar[k]=A+(r<M?r:0)*As0+blockIdx.y*As2; Sr[k]=SFA+(r<M?r:0)*SFAs0+blockIdx.y*SFAs2; }
float ac[4]={0,0,0,0}; int cur=(tid&31)*16;
// Preload all B and SFB into registers
uint4 vb[7]; float sfb_f[7][2];
#pragma unroll
for(int i=0;i<7;++i) {
int off=cur+i*512;
vb[i]=*reinterpret_cast<const uint4*>(s_B+off);
int si=off>>3;
sfb_f[i][0]=s_SFB_f[si]; sfb_f[i][1]=s_SFB_f[si+1];
}
// Main loop
#pragma unroll
for(int i=0;i<7;++i) {
int off=cur+i*512;
int si=off>>3;
uint4 va[4];
#pragma unroll
for(int k=0;k<4;++k) va[k]=load_uint4_cs(Ar[k]+off);
half2 s[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[__ldg(&Sr[k][si*SFAs1])]*sfb_f[i][0]; }
#pragma unroll
for(int k=0;k<4;++k) {
uint32_t ra0, rb0, ra1, rb1, ra2, rb2, ra3, rb3;
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(va[k].x), "r"(vb[i].x));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(va[k].x), "r"(vb[i].x));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(va[k].x), "r"(vb[i].x));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(va[k].x), "r"(vb[i].x));
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), s[k]);
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(va[k].y), "r"(vb[i].y));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(va[k].y), "r"(vb[i].y));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(va[k].y), "r"(vb[i].y));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(va[k].y), "r"(vb[i].y));
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), s[k]);
}
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
#pragma unroll
for(int k=0;k<4;++k) { s[k]=__float2half2_rn(0.0f); sc[k]=f8[__ldg(&Sr[k][(si+1)*SFAs1])]*sfb_f[i][1]; }
#pragma unroll
for(int k=0;k<4;++k) {
uint32_t ra0, rb0, ra1, rb1, ra2, rb2, ra3, rb3;
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(va[k].z), "r"(vb[i].z));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(va[k].z), "r"(vb[i].z));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(va[k].z), "r"(vb[i].z));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(va[k].z), "r"(vb[i].z));
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), s[k]);
asm volatile("{ .reg .b8 ba,bb; mov.b32 {ba,_,_,_}, %2; mov.b32 {bb,_,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra0), "=r"(rb0) : "r"(va[k].w), "r"(vb[i].w));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,ba,_,_}, %2; mov.b32 {_,bb,_,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra1), "=r"(rb1) : "r"(va[k].w), "r"(vb[i].w));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,ba,_}, %2; mov.b32 {_,_,bb,_}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra2), "=r"(rb2) : "r"(va[k].w), "r"(vb[i].w));
asm volatile("{ .reg .b8 ba,bb; mov.b32 {_,_,_,ba}, %2; mov.b32 {_,_,_,bb}, %3; cvt.rn.f16x2.e2m1x2 %0, ba; cvt.rn.f16x2.e2m1x2 %1, bb; }" : "=r"(ra3), "=r"(rb3) : "r"(va[k].w), "r"(vb[i].w));
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra0), *reinterpret_cast<half2*>(&rb0), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra1), *reinterpret_cast<half2*>(&rb1), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra2), *reinterpret_cast<half2*>(&rb2), s[k]);
s[k] = __hfma2(*reinterpret_cast<half2*>(&ra3), *reinterpret_cast<half2*>(&rb3), s[k]);
}
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(s[k].x,s[k].y))*sc[k];
}
#pragma unroll
for(int o=16;o>0;o/=2) for(int k=0;k<4;++k) ac[k]+=__shfl_down_sync(0xffffffff,ac[k],o);
if((tid&31)==0) for(int k=0;k<4;++k) { int r=r0+k; if(r<M) C[r*Cs0+blockIdx.y*Cs2]=__float2half(ac[k]); }
}
// k16384: 512 threads, precomputed f8 LUT + SFB floats
extern "C" __global__ __launch_bounds__(512) void k16384(const uint8_t* __restrict__ A, const uint8_t* __restrict__ B, const uint8_t* __restrict__ SFA, const uint8_t* __restrict__ SFB, half* __restrict__ C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2) {
int tid=threadIdx.x;
extern __shared__ uint8_t sm[];
float* f8=(float*)sm;
uint8_t* s_B=(uint8_t*)(f8+256);
float* s_SFB_f=(float*)(s_B+8192); // SFB as floats (1024 values = 4096 bytes)
// Copy precomputed LUT to shared memory
if(tid<256) f8[tid] = fp8_lut_const[tid];
const uint8_t* B_src=B+blockIdx.y*Bs2;
*reinterpret_cast<uint4*>(s_B+tid*16)=*reinterpret_cast<const uint4*>(B_src+tid*16);
// Convert SFB to floats (1024 scale factors)
const uint8_t* SFB_src=SFB+blockIdx.y*SFBs2;
if(tid<512) s_SFB_f[tid] = fp8_lut_const[SFB_src[tid]];
if(tid<512) s_SFB_f[tid+512] = fp8_lut_const[SFB_src[tid+512]];
__syncthreads();
int r0=blockIdx.x*64+(tid>>5)*4; if(r0>=M) return;
const uint8_t* Ar[4]; const uint8_t* Sr[4];
#pragma unroll
for(int k=0;k<4;++k) { int r=r0+k; if(r<M) { Ar[k]=A+r*As0+blockIdx.y*As2; Sr[k]=SFA+r*SFAs0+blockIdx.y*SFAs2; } else { Ar[k]=Ar[0]; Sr[k]=Sr[0]; } }
float ac[4]={0,0,0,0}; int cur=(tid&31)*16;
uint4 b=*reinterpret_cast<const uint4*>(s_B+cur); uint4 va[4];
#pragma unroll
for(int k=0;k<4;++k) va[k]=load_uint4_cs(Ar[k]+cur);
int si=cur>>3;
float sfb_f[2]; sfb_f[0]=s_SFB_f[si]; sfb_f[1]=s_SFB_f[si+1];
#pragma unroll
for(int s=0;s<15;++s) {
int no=cur+(s+1)*512; uint4 bn=*reinterpret_cast<const uint4*>(s_B+no); uint4 vn[4];
#pragma unroll
for(int k=0;k<4;++k) vn[k]=load_uint4_cs(Ar[k]+no);
int sin=no>>3;
float sfb_fn[2]; sfb_fn[0]=s_SFB_f[sin]; sfb_fn[1]=s_SFB_f[sin+1];
// Compute with current data
half2 sh[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { sh[k]=__float2half2_rn(0.0f); sc[k]=f8[Sr[k][si*SFAs1]]*sfb_f[0]; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(sh[k],va[k].x,b.x); proc_ptx_orig(sh[k],va[k].y,b.y); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(sh[k].x,sh[k].y))*sc[k];
#pragma unroll
for(int k=0;k<4;++k) { sh[k]=__float2half2_rn(0.0f); sc[k]=f8[Sr[k][(si+1)*SFAs1]]*sfb_f[1]; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(sh[k],va[k].z,b.z); proc_ptx_orig(sh[k],va[k].w,b.w); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(sh[k].x,sh[k].y))*sc[k];
// Move to next
b=bn; si=sin; sfb_f[0]=sfb_fn[0]; sfb_f[1]=sfb_fn[1];
#pragma unroll
for(int k=0;k<4;++k) va[k]=vn[k];
}
// Final iteration
half2 sh[4]; float sc[4];
#pragma unroll
for(int k=0;k<4;++k) { sh[k]=__float2half2_rn(0.0f); sc[k]=f8[Sr[k][si*SFAs1]]*sfb_f[0]; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(sh[k],va[k].x,b.x); proc_ptx_orig(sh[k],va[k].y,b.y); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(sh[k].x,sh[k].y))*sc[k];
#pragma unroll
for(int k=0;k<4;++k) { sh[k]=__float2half2_rn(0.0f); sc[k]=f8[Sr[k][(si+1)*SFAs1]]*sfb_f[1]; }
#pragma unroll
for(int k=0;k<4;++k) { proc_ptx_orig(sh[k],va[k].z,b.z); proc_ptx_orig(sh[k],va[k].w,b.w); }
#pragma unroll
for(int k=0;k<4;++k) ac[k]+=__half2float(__hadd(sh[k].x,sh[k].y))*sc[k];
#pragma unroll
for(int o=16;o>0;o/=2) for(int k=0;k<4;++k) ac[k]+=__shfl_down_sync(0xffffffff,ac[k],o);
if((tid&31)==0) for(int k=0;k<4;++k) { int r=r0+k; if(r<M) C[r*Cs0+blockIdx.y*Cs2]=__float2half(ac[k]); }
}
extern "C" __global__ __launch_bounds__(128) void generic(const uint8_t* __restrict__ A, const uint8_t* __restrict__ B, const uint8_t* __restrict__ SFA, const uint8_t* __restrict__ SFB, half* __restrict__ C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2) {
int tid=threadIdx.x; extern __shared__ uint8_t sm[]; uint32_t* l=(uint32_t*)sm; float* f8=(float*)(l+256); uint8_t* ssfb=(uint8_t*)(f8+256);
if(tid<128) { l[tid]=fp4_packed_lut_const[tid]; l[tid+128]=fp4_packed_lut_const[tid+128]; f8[tid]=fp8e4m3_to_float((uint8_t)tid); f8[tid+128]=fp8e4m3_to_float((uint8_t)(tid+128)); }
const uint8_t* Bs=B+blockIdx.y*Bs2; const uint8_t* Ss=SFB+blockIdx.y*SFBs2;
int sfb_sz=K/16; for(int i=tid;i<sfb_sz;i+=128) ssfb[i]=Ss[i];
__syncthreads();
int r0=blockIdx.x*16+(tid>>5)*4; if(r0>=M) return;
const uint8_t* Ar[4]; const uint8_t* Sr[4];
#pragma unroll
for(int k=0;k<4;++k) { int r=r0+k; Ar[k]=A+(r<M?r:0)*As0+blockIdx.y*As2; Sr[k]=SFA+(r<M?r:0)*SFAs0+blockIdx.y*SFAs2; }
float ac[4]={0,0,0,0}; int b_bytes=K/2; int cur=(tid&31)*16;
while(cur<b_bytes) {
uint4 b=load_uint4_ca(Bs+cur); uint4 va[4];
#pragma unroll
for(int k=0;k<4;++k) va[k]=load_uint4_cs(Ar[k]+cur);
uint8_t sfa[4][2], sfb[2]; int si=cur>>3;
#pragma unroll
for(int k=0;k<4;++k) { sfa[k][0]=Sr[k][si*SFAs1]; sfa[k][1]=Sr[k][(si+1)*SFAs1]; }
sfb[0]=ssfb[si]; sfb[1]=ssfb[si+1];
blk_lut(ac,va,b,sfa,sfb,l,f8);
cur+=512;
}
#pragma unroll
for(int o=16;o>0;o/=2) for(int k=0;k<4;++k) ac[k]+=__shfl_down_sync(0xffffffff,ac[k],o);
if((tid&31)==0) for(int k=0;k<4;++k) { int r=r0+k; if(r<M) C[r*Cs0+blockIdx.y*Cs2]=__float2half(ac[k]); }
}
extern "C" void launch_nvfp4_gemv_fast_v7(long long A, long long B, long long SFA, long long SFB, long long C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2) {
if (K==2048) {
// Shared mem: 1024 (f8) + 1024 (B) + 128 (SFB bytes) = 2176
k2048<<<dim3((M+31)/32,L),dim3(256),1024+1024+128>>>((uint8_t*)A,(uint8_t*)B,(uint8_t*)SFA,(uint8_t*)SFB,(half*)C,M,K,L,As0,As1,As2,Bs0,Bs1,Bs2,SFAs0,SFAs1,SFAs2,SFBs0,SFBs1,SFBs2,Cs0,Cs1,Cs2);
} else if (K==7168) {
// 128 threads, 16 rows/block (4 rows/warp) - optimal config
// Shared mem: 1024 (f8) + 3584 (B) + 1792 (SFB floats 448*4) = 6400
k7168<<<dim3((M+15)/16,L),dim3(128),1024+3584+1792>>>((uint8_t*)A,(uint8_t*)B,(uint8_t*)SFA,(uint8_t*)SFB,(half*)C,M,K,L,As0,As1,As2,Bs0,Bs1,Bs2,SFAs0,SFAs1,SFAs2,SFBs0,SFBs1,SFBs2,Cs0,Cs1,Cs2);
} else if (K==16384) {
// 512 threads, 64 rows/block - Shared mem: 1024 (f8) + 8192 (B) + 4096 (SFB floats) = 13312
k16384<<<dim3((M+63)/64,L),dim3(512),1024+8192+4096>>>((uint8_t*)A,(uint8_t*)B,(uint8_t*)SFA,(uint8_t*)SFB,(half*)C,M,K,L,As0,As1,As2,Bs0,Bs1,Bs2,SFAs0,SFAs1,SFAs2,SFBs0,SFBs1,SFBs2,Cs0,Cs1,Cs2);
} else {
generic<<<dim3((M+15)/16,L),dim3(128),1024+4096>>>((uint8_t*)A,(uint8_t*)B,(uint8_t*)SFA,(uint8_t*)SFB,(half*)C,M,K,L,As0,As1,As2,Bs0,Bs1,Bs2,SFAs0,SFAs1,SFAs2,SFBs0,SFBs1,SFBs2,Cs0,Cs1,Cs2);
}
}
"""
cpp_src = """
extern "C" void launch_nvfp4_gemv_fast_v7(long long A, long long B, long long SFA, long long SFB, long long C, int M, int K, int L, long long As0, long long As1, long long As2, long long Bs0, long long Bs1, long long Bs2, long long SFAs0, long long SFAs1, long long SFAs2, long long SFBs0, long long SFBs1, long long SFBs2, long long Cs0, long long Cs1, long long Cs2);
"""
import torch.utils.cpp_extension
module = torch.utils.cpp_extension.load_inline(
name="fp4_gemv_k16384_final",
cpp_sources=cpp_src,
cuda_sources=CUDA_SOURCE,
functions=["launch_nvfp4_gemv_fast_v7"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-gencode=arch=compute_100a,code=sm_100a", "--maxrregcount=48"],
verbose=True,
)
def custom_kernel(input: input_t) -> output_t:
A, B, SFA, SFB, _, _, C = input
M, K_packed, L = A.shape
K = K_packed * 2
module.launch_nvfp4_gemv_fast_v7(A.data_ptr(), B.data_ptr(), SFA.data_ptr(), SFB.data_ptr(), C.data_ptr(), M, K, L, A.stride(0), A.stride(1), A.stride(2), B.stride(0), B.stride(1), B.stride(2), SFA.stride(0), SFA.stride(1), SFA.stride(2), SFB.stride(0), SFB.stride(1), SFB.stride(2), C.stride(0), C.stride(1), C.stride(2))
return C
scrolls · 535 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON