impl Default for CudaKernels {
fn default() -> Self {
Self::new()
}
}
fn generate_q4_0_candle_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// BUG-GGUF-001 FIX: Q4_0 GEMV with candle nibble layout
// Each warp (32 threads) computes one output element
// Thread 0-15: use low nibbles from bytes 0-15
// Thread 16-31: use high nibbles from bytes 0-15
.visible .entry q4_0_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<32>;
.reg .u64 %rd<16>;
.reg .f32 %f<16>;
.reg .b16 %h<4>;
.reg .pred %p<8>;
// r0=tid, r1=ctaid, r2=n_dim, r3=k_dim
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
// Bounds check: if ctaid >= n_dim, exit
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_exit;
// f0 = accumulator
mov.f32 %f0, 0f00000000;
// r4 = num_blocks = ceil(k_dim / 32)
add.u32 %r4, %r3, 31;
shr.u32 %r4, %r4, 5;
// rd3 = row_base = w_ptr + ctaid * num_blocks * 18
mul.lo.u32 %r5, %r4, 18;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// r6 = blk_idx (loop counter)
mov.u32 %r6, 0;
$L_blk_loop:
setp.ge.u32 %p1, %r6, %r4;
@%p1 bra $L_blk_loop_end;
// rd4 = blk_addr = row_base + blk_idx * 18
mul.wide.u32 %rd4, %r6, 18;
add.u64 %rd4, %rd3, %rd4;
// f1 = scale d (fp16 at offset 0) - use b16 register for f16 conversion
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// rd5 = qs_base = blk_addr + 2
add.u64 %rd5, %rd4, 2;
// CANDLE LAYOUT:
// Thread 0-15 read bytes 0-15 (low nibbles -> positions 0-15)
// Thread 16-31 read bytes 0-15 (high nibbles -> positions 16-31)
// r8 = byte_idx = tid < 16 ? tid : tid - 16
setp.ge.u32 %p2, %r0, 16;
mov.u32 %r8, %r0;
@%p2 sub.u32 %r8, %r0, 16;
// Load byte from qs[byte_idx]
cvt.u64.u32 %rd6, %r8;
add.u64 %rd6, %rd5, %rd6;
ld.global.u8 %r9, [%rd6];
// r10 = nibble value
// Threads 0-15: low nibble (byte & 0xF)
// Threads 16-31: high nibble (byte >> 4)
mov.u32 %r10, %r9;
@%p2 shr.u32 %r10, %r9, 4;
and.b32 %r10, %r10, 15;
// r11 = centered value = nibble - 8 (as signed)
sub.u32 %r11, %r10, 8;
// f2 = dequantized = d * centered
cvt.rn.f32.s32 %f2, %r11;
mul.f32 %f2, %f1, %f2;
// r12 = x_idx = blk_idx * 32 + tid
shl.b32 %r12, %r6, 5;
add.u32 %r12, %r12, %r0;
// Bounds check for last block
setp.ge.u32 %p3, %r12, %r3;
@%p3 bra $L_skip_mul;
// f3 = x[x_idx]
cvt.u64.u32 %rd7, %r12;
shl.b64 %rd7, %rd7, 2;
add.u64 %rd7, %rd2, %rd7;
ld.global.f32 %f3, [%rd7];
// f0 += f2 * f3
fma.rn.f32 %f0, %f2, %f3, %f0;
$L_skip_mul:
add.u32 %r6, %r6, 1;
bra $L_blk_loop;
$L_blk_loop_end:
// Warp reduction using shfl.sync.down
shfl.sync.down.b32 %f4, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f4;
shfl.sync.down.b32 %f5, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f5;
shfl.sync.down.b32 %f6, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f6;
shfl.sync.down.b32 %f7, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f7;
shfl.sync.down.b32 %f8, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f8;
// Thread 0 writes result
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_exit;
// y[ctaid] = f0
mul.wide.u32 %rd8, %r1, 4;
add.u64 %rd8, %rd0, %rd8;
st.global.f32 [%rd8], %f0;
$L_exit:
ret;
}
",
)
}
fn generate_q4_1_candle_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// PMAT-782 FIX: Q4_1 GEMV with candle nibble layout (affine dequant d*q + m)
// Each warp (32 threads) computes one output element
// Thread 0-15: use low nibbles from bytes 0-15
// Thread 16-31: use high nibbles from bytes 0-15
.visible .entry q4_1_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<32>;
.reg .u64 %rd<16>;
.reg .f32 %f<16>;
.reg .b16 %h<4>;
.reg .pred %p<8>;
// r0=tid, r1=ctaid, r2=n_dim, r3=k_dim
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
// Bounds check: if ctaid >= n_dim, exit
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_exit;
// f0 = accumulator
mov.f32 %f0, 0f00000000;
// r4 = num_blocks = ceil(k_dim / 32)
add.u32 %r4, %r3, 31;
shr.u32 %r4, %r4, 5;
// rd3 = row_base = w_ptr + ctaid * num_blocks * 20
mul.lo.u32 %r5, %r4, 20;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// r6 = blk_idx (loop counter)
mov.u32 %r6, 0;
$L_blk_loop:
setp.ge.u32 %p1, %r6, %r4;
@%p1 bra $L_blk_loop_end;
// rd4 = blk_addr = row_base + blk_idx * 20
mul.wide.u32 %rd4, %r6, 20;
add.u64 %rd4, %rd3, %rd4;
// f1 = scale d (fp16 at offset 0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// f4 = min m (fp16 at offset 2)
add.u64 %rd5, %rd4, 2;
ld.global.b16 %h1, [%rd5];
cvt.f32.f16 %f4, %h1;
// rd6 = qs_base = blk_addr + 4
add.u64 %rd6, %rd4, 4;
// CANDLE LAYOUT:
// Thread 0-15 read bytes 0-15 (low nibbles -> positions 0-15)
// Thread 16-31 read bytes 0-15 (high nibbles -> positions 16-31)
// r8 = byte_idx = tid < 16 ? tid : tid - 16
setp.ge.u32 %p2, %r0, 16;
mov.u32 %r8, %r0;
@%p2 sub.u32 %r8, %r0, 16;
// Load byte from qs[byte_idx]
cvt.u64.u32 %rd7, %r8;
add.u64 %rd7, %rd6, %rd7;
ld.global.u8 %r9, [%rd7];
// r10 = nibble value
// Threads 0-15: low nibble (byte & 0xF)
// Threads 16-31: high nibble (byte >> 4)
mov.u32 %r10, %r9;
@%p2 shr.u32 %r10, %r9, 4;
and.b32 %r10, %r10, 15;
// Q4_1 affine dequant: f2 = d * nibble + m (no centering)
cvt.rn.f32.u32 %f2, %r10;
fma.rn.f32 %f2, %f1, %f2, %f4;
// r12 = x_idx = blk_idx * 32 + tid
shl.b32 %r12, %r6, 5;
add.u32 %r12, %r12, %r0;
// Bounds check for last block
setp.ge.u32 %p3, %r12, %r3;
@%p3 bra $L_skip_mul;
// f3 = x[x_idx]
cvt.u64.u32 %rd8, %r12;
shl.b64 %rd8, %rd8, 2;
add.u64 %rd8, %rd2, %rd8;
ld.global.f32 %f3, [%rd8];
// f0 += f2 * f3
fma.rn.f32 %f0, %f2, %f3, %f0;
$L_skip_mul:
add.u32 %r6, %r6, 1;
bra $L_blk_loop;
$L_blk_loop_end:
// Warp reduction using shfl.sync.down
shfl.sync.down.b32 %f5, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f5;
shfl.sync.down.b32 %f6, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f6;
shfl.sync.down.b32 %f7, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f7;
shfl.sync.down.b32 %f8, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f8;
shfl.sync.down.b32 %f9, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f9;
// Thread 0 writes result
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_exit;
// y[ctaid] = f0
mul.wide.u32 %rd9, %r1, 4;
add.u64 %rd9, %rd0, %rd9;
st.global.f32 [%rd9], %f0;
$L_exit:
ret;
}
",
)
}
fn generate_iq4_xs_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// The 16 non-linear IQ4_NL levels: quantize::iq_grids::KVALUES_IQ4NL.
.global .align 4 .s32 kvalues_iq4nl[16] = {-127, -104, -83, -65, -49, -35, -22, -10, 1, 13, 25, 38, 53, 69, 89, 113};
.visible .entry iq4_xs_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<48>;
.reg .u64 %rd<32>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_iq_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 136
mul.lo.u32 %r5, %r4, 136;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// per-thread: jhalf = tid >> 4 (nibble select), jlow = tid & 15 (byte index)
shr.u32 %r6, %r0, 4;
and.b32 %r7, %r0, 15;
// codebook base
mov.u64 %rd10, kvalues_iq4nl;
mov.u32 %r8, 0; // blk
$L_iq_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_iq_blk_end;
// blk_addr = row_base + blk * 136
mul.wide.u32 %rd4, %r8, 136;
add.u64 %rd4, %rd3, %rd4;
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// scales_h (u16 at +2)
add.u64 %rd5, %rd4, 2;
ld.global.u16 %r9, [%rd5];
mov.u32 %r10, 0; // m = sub-block index = ib
$L_iq_sub:
setp.ge.u32 %p2, %r10, 8;
@%p2 bra $L_iq_sub_end;
// ls_low = (scales_l[m >> 1] >> (4 * (m & 1))) & 0xf
shr.u32 %r11, %r10, 1;
cvt.u64.u32 %rd6, %r11;
add.u64 %rd6, %rd4, %rd6;
add.u64 %rd6, %rd6, 4;
ld.global.u8 %r12, [%rd6];
and.b32 %r13, %r10, 1;
shl.b32 %r13, %r13, 2;
shr.u32 %r12, %r12, %r13;
and.b32 %r12, %r12, 15;
// ls_high = ((scales_h >> (2 * m)) & 3) << 4
shl.b32 %r14, %r10, 1;
shr.u32 %r15, %r9, %r14;
and.b32 %r15, %r15, 3;
shl.b32 %r15, %r15, 4;
// ls = ls_low | ls_high ; dl = d * (ls - 32)
or.b32 %r16, %r12, %r15;
cvt.rn.f32.u32 %f2, %r16;
sub.f32 %f2, %f2, 0f42000000; // 32.0
mul.f32 %f3, %f1, %f2;
// byte = qs[16*m + jlow] (qs starts at +8)
shl.b32 %r17, %r10, 4;
add.u32 %r17, %r17, %r7;
cvt.u64.u32 %rd7, %r17;
add.u64 %rd7, %rd4, %rd7;
add.u64 %rd7, %rd7, 8;
ld.global.u8 %r18, [%rd7];
// nib = jhalf ? (byte >> 4) : (byte & 0xf)
shl.b32 %r19, %r6, 2; // jhalf * 4
shr.u32 %r20, %r18, %r19;
and.b32 %r20, %r20, 15;
// w = dl * kvalues_iq4nl[nib]
mul.wide.u32 %rd8, %r20, 4;
add.u64 %rd8, %rd10, %rd8;
ld.global.s32 %r21, [%rd8];
cvt.rn.f32.s32 %f4, %r21;
mul.f32 %f5, %f3, %f4;
// x_idx = blk*256 + m*32 + tid
shl.b32 %r22, %r8, 8;
shl.b32 %r23, %r10, 5;
add.u32 %r22, %r22, %r23;
add.u32 %r22, %r22, %r0;
setp.ge.u32 %p3, %r22, %r3;
@%p3 bra $L_iq_skip;
mul.wide.u32 %rd9, %r22, 4;
add.u64 %rd9, %rd2, %rd9;
ld.global.f32 %f6, [%rd9];
fma.rn.f32 %f0, %f5, %f6, %f0;
$L_iq_skip:
add.u32 %r10, %r10, 1;
bra $L_iq_sub;
$L_iq_sub_end:
add.u32 %r8, %r8, 1;
bra $L_iq_blk;
$L_iq_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_iq_exit;
mul.wide.u32 %rd11, %r1, 4;
add.u64 %rd11, %rd0, %rd11;
st.global.f32 [%rd11], %f0;
$L_iq_exit:
ret;
}
",
)
}
fn generate_q5_1_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
.visible .entry q5_1_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<40>;
.reg .u64 %rd<24>;
.reg .f32 %f<20>;
.reg .b16 %h<4>;
.reg .pred %p<10>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_q51_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 32)
add.u32 %r4, %r3, 31;
shr.u32 %r4, %r4, 5;
// row_base = w_ptr + ctaid * nb * 24
mul.lo.u32 %r5, %r4, 24;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// jlow = tid & 15 (byte index), jhalf = tid >> 4 (nibble select)
and.b32 %r7, %r0, 15;
shr.u32 %r6, %r0, 4;
// nibble shift = jhalf * 4
shl.b32 %r19, %r6, 2;
// qh bit index = jlow + 16*jhalf
shl.b32 %r23, %r6, 4;
add.u32 %r23, %r23, %r7;
mov.u32 %r8, 0;
$L_q51_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_q51_blk_end;
// blk_addr = row_base + blk * 24
mul.wide.u32 %rd4, %r8, 24;
add.u64 %rd4, %rd3, %rd4;
// d (f16 at +0), m (f16 at +2)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
ld.global.b16 %h1, [%rd4+2];
cvt.f32.f16 %f2, %h1;
// qh (u32 at +4)
ld.global.u32 %r24, [%rd4+4];
// byte = qs[jlow], qs starts at +8
cvt.u64.u32 %rd7, %r7;
add.u64 %rd7, %rd4, %rd7;
add.u64 %rd7, %rd7, 8;
ld.global.u8 %r18, [%rd7];
// nib = (byte >> (4*jhalf)) & 0xf
shr.u32 %r20, %r18, %r19;
and.b32 %r20, %r20, 15;
// hb = (qh >> (jlow + 16*jhalf)) & 1, placed at bit 4
shr.u32 %r25, %r24, %r23;
and.b32 %r25, %r25, 1;
shl.b32 %r25, %r25, 4;
or.b32 %r20, %r20, %r25;
// w = q * d + m
cvt.rn.f32.u32 %f4, %r20;
fma.rn.f32 %f5, %f4, %f1, %f2;
// x_idx = blk*32 + tid
shl.b32 %r22, %r8, 5;
add.u32 %r22, %r22, %r0;
setp.ge.u32 %p3, %r22, %r3;
@%p3 bra $L_q51_skip;
mul.wide.u32 %rd9, %r22, 4;
add.u64 %rd9, %rd2, %rd9;
ld.global.f32 %f6, [%rd9];
fma.rn.f32 %f0, %f5, %f6, %f0;
$L_q51_skip:
add.u32 %r8, %r8, 1;
bra $L_q51_blk;
$L_q51_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_q51_exit;
mul.wide.u32 %rd11, %r1, 4;
add.u64 %rd11, %rd0, %rd11;
st.global.f32 [%rd11], %f0;
$L_q51_exit:
ret;
}
",
)
}
fn generate_iq3_s_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// IQ3S_GRID: quantize::iq_grids::IQ3S_GRID, 512 packed 4-byte codebook entries.
.global .align 4 .u32 iq3s_grid_g[512] = {
16843009, 16843011, 16843013, 16843019, 16843023, 16843521, 16843523, 16843525,
16843529, 16843533, 16844033, 16844035, 16844043, 16844551, 16845057, 16845061,
16845067, 16845071, 16845571, 16845575, 16846081, 16846085, 16846595, 16846601,
16846607, 16974081, 16974083, 16974085, 16974089, 16974593, 16974595, 16974603,
16975105, 16975111, 16975119, 16975619, 16975627, 16976137, 16977155, 16977163,
16977669, 17105153, 17105155, 17105163, 17105167, 17105665, 17105671, 17105677,
17106179, 17106187, 17106689, 17106697, 17107205, 17107211, 17107215, 17107715,
17107719, 17108737, 17108743, 17236231, 17236739, 17236747, 17237249, 17237253,
17237763, 17237767, 17237773, 17238281, 17238785, 17238789, 17239311, 17239811,
17239819, 17367297, 17367815, 17367823, 17368323, 17368329, 17368837, 17369345,
17369351, 17369859, 17370881, 17498373, 17498377, 17499393, 17499397, 17499405,
17499911, 17500419, 17500427, 17500431, 17501453, 17501959, 17629453, 17629955,
17629959, 17630979, 17632005, 17633027, 17760513, 17760517, 17760521, 17761537,
17761541, 17761549, 17762055, 17763073, 17763081, 50397441, 50397443, 50397445,
50397449, 50397953, 50397955, 50397959, 50397963, 50397967, 50398465, 50398469,
50398979, 50398985, 50398989, 50400009, 50400013, 50400515, 50401029, 50528513,
50528515, 50528519, 50528525, 50529025, 50529033, 50529539, 50530049, 50530055,
50530563, 50531073, 50531077, 50532097, 50532109, 50659585, 50660101, 50660107,
50660111, 50660609, 50660617, 50661125, 50661633, 50661639, 50662155, 50662657,
50663173, 50790659, 50790665, 50790671, 50791169, 50791175, 50791683, 50791695,
50792193, 50792201, 50792707, 50793733, 50794241, 50921735, 50921739, 50922245,
50922249, 50923267, 50923271, 50923781, 50923789, 50924289, 50924297, 51052803,
51053313, 51053319, 51053827, 51054337, 51054341, 51055363, 51184897, 51184905,
51184911, 51185929, 51185933, 51314947, 51314951, 51315457, 51315461, 51315971,
51316491, 51316995, 51318021, 51318529, 83951873, 83951875, 83951879, 83951883,
83951887, 83952385, 83952389, 83952393, 83952397, 83952899, 83952903, 83952911,
83953409, 83953413, 83953923, 83953927, 83953931, 83954433, 83954437, 83954959,
83955457, 83955463, 83955467, 84082945, 84082949, 84083457, 84083463, 84083471,
84083973, 84083979, 84084483, 84084489, 84084997, 84085507, 84214019, 84214025,
84214031, 84215043, 84215047, 84215553, 84215567, 84216067, 84216583, 84216591,
84217603, 84217609, 84345089, 84345093, 84345099, 84345603, 84346117, 84346121,
84346627, 84346631, 84347141, 84347649, 84348173, 84476163, 84476175, 84477185,
84477191, 84477701, 84477707, 84478211, 84479749, 84479755, 84607241, 84607747,
84608261, 84608783, 84609281, 84609799, 84610817, 84738305, 84738309, 84738319,
84739331, 84740875, 84741379, 84869387, 84869891, 84870413, 84870913, 84871431,
84871937, 117506309, 117506819, 117506823, 117506827, 117506831, 117507333, 117507843,
117507847, 117507851, 117508357, 117508361, 117508367, 117508867, 117509383, 117509891,
117637379, 117637383, 117637387, 117637897, 117638403, 117638407, 117639425, 117640449,
117640965, 117640973, 117768449, 117768965, 117769473, 117769989, 117769993, 117771009,
117899523, 117900033, 117900041, 117900547, 117900551, 117900559, 117901057, 117901571,
117901575, 117901583, 117902091, 117903111, 118030599, 118031107, 118031117, 118031621,
118032131, 118033157, 118033665, 118033673, 118161667, 118162177, 118162181, 118162699,
118163205, 118163721, 118164237, 118165255, 118293261, 118294787, 118423811, 118423815,
118424833, 118424837, 118425355, 151060737, 151060745, 151061253, 151061761, 151061769,
151061775, 151062277, 151062787, 151063297, 151064321, 151191813, 151191823, 151192323,
151192327, 151192837, 151193345, 151193355, 151193863, 151194371, 151194379, 151322883,
151322887, 151323393, 151323403, 151323907, 151324423, 151324929, 151325455, 151325957,
151326465, 151453961, 151454467, 151454471, 151454977, 151454981, 151455491, 151455499,
151585025, 151585029, 151586057, 151586575, 151587073, 151588611, 151716107, 151716111,
151717123, 151719173, 151847687, 151848713, 151850241, 151978753, 151978763, 151979777,
151980295, 151980803, 184615173, 184615681, 184615689, 184616197, 184617217, 184617225,
184617231, 184617733, 184618253, 184618761, 184746243, 184746247, 184746251, 184746757,
184747267, 184747781, 184749829, 184877313, 184877827, 184878343, 184878849, 184878861,
184879879, 185008389, 185008399, 185008897, 185009423, 185010441, 185010947, 185011467,
185011975, 185139459, 185139465, 185140481, 185140997, 185141517, 185271045, 185271565,
185273091, 185273095, 185403653, 185532677, 185532681, 185533701, 218170115, 218170119,
218170123, 218171139, 218171143, 218172673, 218300673, 218301697, 218301711, 218303753,
218432261, 218433289, 218433797, 218434315, 218434821, 218435329, 218562817, 218563337,
218563843, 218564865, 218694923, 218695943, 218696965, 218824961, 218824967, 218826505,
218828033, 218956043, 218958081, 219087619, 219087623, 251724033, 251724041, 251724047,
251725057, 251725061, 251725581, 251726081, 251726601, 251727109, 251855109, 251855619,
251856137, 251857159, 251857163, 251986179, 251986185, 251986689, 251986701, 251987203,
251987713, 251988739, 252117253, 252118789, 252118795, 252119815, 252248323, 252248331,
252248839, 252249345, 252250881, 252380421, 252381445, 252510469, 252512003, 252641537
};
.visible .entry iq3_s_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<48>;
.reg .u64 %rd<32>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_s3_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 110
mul.lo.u32 %r5, %r4, 110;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// ib = tid >> 2 (0..8), l = tid & 3 (0..4)
shr.u32 %r6, %r0, 2;
and.b32 %r7, %r0, 3;
mov.u64 %rd10, iq3s_grid_g;
mov.u32 %r8, 0; // blk
$L_s3_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_s3_blk_end;
mul.wide.u32 %rd4, %r8, 110;
add.u64 %rd4, %rd3, %rd4;
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// sc = scales[ib >> 1] (scales at +106)
shr.u32 %r9, %r6, 1;
cvt.u64.u32 %rd5, %r9;
add.u64 %rd5, %rd4, %rd5;
add.u64 %rd5, %rd5, 106;
ld.global.u8 %r10, [%rd5];
// nibble = ib & 1 ? (sc >> 4) : (sc & 0xf)
and.b32 %r11, %r6, 1;
shl.b32 %r12, %r11, 2; // 0 or 4
shr.u32 %r13, %r10, %r12;
and.b32 %r13, %r13, 15;
// db = d * (1 + 2*nib_val)
cvt.rn.f32.u32 %f2, %r13;
add.f32 %f2, %f2, %f2; // 2*v
add.f32 %f2, %f2, 0f3F800000; // +1.0
mul.f32 %f3, %f1, %f2; // db
// qh_byte = qh[ib] (qh at +66)
cvt.u64.u32 %rd6, %r6;
add.u64 %rd6, %rd4, %rd6;
add.u64 %rd6, %rd6, 66;
ld.global.u8 %r14, [%rd6];
// qs base = +2 + 8*ib + 2*l
shl.b32 %r15, %r6, 3;
shl.b32 %r16, %r7, 1;
add.u32 %r15, %r15, %r16;
cvt.u64.u32 %rd7, %r15;
add.u64 %rd7, %rd4, %rd7;
add.u64 %rd7, %rd7, 2;
ld.global.u8 %r17, [%rd7]; // qs[8ib+2l]
ld.global.u8 %r18, [%rd7+1]; // qs[8ib+2l+1]
// high bits: bit(2l) and bit(2l+1) of qh_byte, each placed at 256
shr.u32 %r19, %r14, %r16; // qh >> 2l
and.b32 %r20, %r19, 1;
shl.b32 %r20, %r20, 8;
or.b32 %r17, %r17, %r20; // i1
shr.u32 %r21, %r19, 1; // qh >> (2l+1)
and.b32 %r21, %r21, 1;
shl.b32 %r21, %r21, 8;
or.b32 %r18, %r18, %r21; // i2
// g1 = grid[i1], g2 = grid[i2]
mul.wide.u32 %rd8, %r17, 4;
add.u64 %rd8, %rd10, %rd8;
ld.global.u32 %r22, [%rd8];
mul.wide.u32 %rd9, %r18, 4;
add.u64 %rd9, %rd10, %rd9;
ld.global.u32 %r23, [%rd9];
// sign_byte = signs[4*ib + l] (signs at +74)
shl.b32 %r24, %r6, 2;
add.u32 %r24, %r24, %r7;
cvt.u64.u32 %rd11, %r24;
add.u64 %rd11, %rd4, %rd11;
add.u64 %rd11, %rd11, 74;
ld.global.u8 %r25, [%rd11];
// col0 = blk*256 + 32*ib + 8*l
shl.b32 %r26, %r8, 8;
shl.b32 %r27, %r6, 5;
add.u32 %r26, %r26, %r27;
shl.b32 %r28, %r7, 3;
add.u32 %r26, %r26, %r28;
mov.u32 %r29, 0; // j = 0..4
$L_s3_j:
setp.ge.u32 %p2, %r29, 4;
@%p2 bra $L_s3_j_end;
shl.b32 %r30, %r29, 3; // 8*j
shr.u32 %r31, %r22, %r30;
and.b32 %r31, %r31, 255; // m1
shr.u32 %r32, %r23, %r30;
and.b32 %r32, %r32, 255; // m2
// s1 = bit j of sign_byte, s2 = bit j+4
shr.u32 %r33, %r25, %r29;
and.b32 %r33, %r33, 1;
add.u32 %r34, %r29, 4;
shr.u32 %r35, %r25, %r34;
and.b32 %r35, %r35, 1;
cvt.rn.f32.u32 %f4, %r31;
mul.f32 %f4, %f4, %f3; // db * m1
setp.ne.u32 %p3, %r33, 0;
@%p3 neg.f32 %f4, %f4;
cvt.rn.f32.u32 %f5, %r32;
mul.f32 %f5, %f5, %f3; // db * m2
setp.ne.u32 %p4, %r35, 0;
@%p4 neg.f32 %f5, %f5;
// x[col0 + j]
add.u32 %r36, %r26, %r29;
setp.ge.u32 %p5, %r36, %r3;
@%p5 bra $L_s3_skip1;
mul.wide.u32 %rd12, %r36, 4;
add.u64 %rd12, %rd2, %rd12;
ld.global.f32 %f6, [%rd12];
fma.rn.f32 %f0, %f4, %f6, %f0;
$L_s3_skip1:
// x[col0 + j + 4]
add.u32 %r37, %r36, 4;
setp.ge.u32 %p6, %r37, %r3;
@%p6 bra $L_s3_skip2;
mul.wide.u32 %rd13, %r37, 4;
add.u64 %rd13, %rd2, %rd13;
ld.global.f32 %f7, [%rd13];
fma.rn.f32 %f0, %f5, %f7, %f0;
$L_s3_skip2:
add.u32 %r29, %r29, 1;
bra $L_s3_j;
$L_s3_j_end:
add.u32 %r8, %r8, 1;
bra $L_s3_blk;
$L_s3_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p7, %r0, 0;
@%p7 bra $L_s3_exit;
mul.wide.u32 %rd14, %r1, 4;
add.u64 %rd14, %rd0, %rd14;
st.global.f32 [%rd14], %f0;
$L_s3_exit:
ret;
}
",
)
}
fn generate_q2_k_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
.visible .entry q2_k_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<40>;
.reg .u64 %rd<24>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_q2_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 84
mul.lo.u32 %r5, %r4, 84;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// lane decomposition: g = L>>4, s = (L&15)>>2, h = (L&3)>>1, odd = L&1
shr.u32 %r6, %r0, 4; // g
and.b32 %r7, %r0, 15;
shr.u32 %r7, %r7, 2; // s
and.b32 %r8, %r0, 3;
shr.u32 %r8, %r8, 1; // h
and.b32 %r9, %r0, 1; // odd
// scale index = 8g + 2s + h
shl.b32 %r10, %r6, 3;
shl.b32 %r11, %r7, 1;
add.u32 %r10, %r10, %r11;
add.u32 %r10, %r10, %r8; // sidx
// shift = 2s
shl.b32 %r12, %r7, 1; // shift
// qs run offset = 16 + 32g + 16h + 8*odd
shl.b32 %r13, %r6, 5;
shl.b32 %r14, %r8, 4;
add.u32 %r13, %r13, %r14;
shl.b32 %r14, %r9, 3;
add.u32 %r13, %r13, %r14;
add.u32 %r13, %r13, 16; // qoff
// this lane's first output within a block = 8L
shl.b32 %r15, %r0, 3;
cvt.u64.u32 %rd4, %r10; // sidx
cvt.u64.u32 %rd5, %r13; // qoff
mov.u32 %r16, 0; // blk
$L_q2_blk:
setp.ge.u32 %p1, %r16, %r4;
@%p1 bra $L_q2_blk_end;
mul.wide.u32 %rd6, %r16, 84;
add.u64 %rd6, %rd3, %rd6; // block base
// d (f16 at +80), dmin (f16 at +82)
ld.global.b16 %h0, [%rd6+80];
cvt.f32.f16 %f1, %h0; // d
ld.global.b16 %h1, [%rd6+82];
cvt.f32.f16 %f2, %h1; // dmin
// sc = scales[sidx]: low nibble scales d, high nibble scales dmin
add.u64 %rd7, %rd6, %rd4;
ld.global.u8 %r17, [%rd7];
and.b32 %r18, %r17, 15; // scale nibble
shr.u32 %r19, %r17, 4; // min nibble
cvt.rn.f32.u32 %f3, %r18;
mul.rn.f32 %f3, %f1, %f3; // dl = d * lo
cvt.rn.f32.u32 %f4, %r19;
mul.rn.f32 %f4, %f2, %f4; // ml = dmin * hi
add.u64 %rd8, %rd6, %rd5; // this lane's qs run
// col0 = blk*256 + 8L
shl.b32 %r20, %r16, 8;
add.u32 %r20, %r20, %r15;
mov.u32 %r21, 0; // j = 0..8
$L_q2_j:
setp.ge.u32 %p2, %r21, 8;
@%p2 bra $L_q2_j_end;
cvt.u64.u32 %rd9, %r21;
add.u64 %rd9, %rd8, %rd9;
ld.global.u8 %r22, [%rd9];
shr.u32 %r22, %r22, %r12;
and.b32 %r22, %r22, 3; // q
cvt.rn.f32.u32 %f5, %r22;
mul.rn.f32 %f5, %f3, %f5; // dl * q
sub.rn.f32 %f5, %f5, %f4; // - ml (the affine min)
add.u32 %r23, %r20, %r21; // col
setp.ge.u32 %p3, %r23, %r3;
@%p3 bra $L_q2_skip;
mul.wide.u32 %rd10, %r23, 4;
add.u64 %rd10, %rd2, %rd10;
ld.global.f32 %f6, [%rd10];
fma.rn.f32 %f0, %f5, %f6, %f0;
$L_q2_skip:
add.u32 %r21, %r21, 1;
bra $L_q2_j;
$L_q2_j_end:
add.u32 %r16, %r16, 1;
bra $L_q2_blk;
$L_q2_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p7, %r0, 0;
@%p7 bra $L_q2_exit;
mul.wide.u32 %rd14, %r1, 4;
add.u64 %rd14, %rd0, %rd14;
st.global.f32 [%rd14], %f0;
$L_q2_exit:
ret;
}
",
)
}
fn generate_iq2_xxs_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
let grid = crate::quantize::iq_grids::IQ2XXS_GRID
.iter()
.flat_map(|&v| [(v & 0xffff_ffff) as u32, (v >> 32) as u32])
.map(|w| w.to_string())
.collect::<Vec<_>>()
.join(", ");
let signs = crate::quantize::iq_grids::KSIGNS_IQ2XS
.iter()
.map(u8::to_string)
.collect::<Vec<_>>()
.join(", ");
let mut ptx = String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// IQ2XXS_GRID: 256 u64 codebook entries as 512 u32 (lo, hi). GENERATED from
// quantize::iq_grids::IQ2XXS_GRID when this PTX is built, never hand-copied.
.global .align 4 .u32 iq2xxs_grid_g[512] = {",
);
ptx.push_str(&grid);
ptx.push_str(
r"};
// KSIGNS_IQ2XS: 128 sign bytes. GENERATED from quantize::iq_grids::KSIGNS_IQ2XS.
.global .align 1 .u8 iq2xs_ksigns_g[128] = {",
);
ptx.push_str(&signs);
ptx.push_str(
r"};
.visible .entry iq2_xxs_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<48>;
.reg .u64 %rd<32>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_x2_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 66
mul.lo.u32 %r5, %r4, 66;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// ib = tid >> 2 (0..8), l = tid & 3 (0..4)
shr.u32 %r6, %r0, 2;
and.b32 %r7, %r0, 3;
mov.u64 %rd10, iq2xxs_grid_g;
mov.u64 %rd15, iq2xs_ksigns_g;
// this thread's sub-block offset in the block: 2 + 8*ib
shl.b32 %r9, %r6, 3;
add.u32 %r9, %r9, 2;
// this thread's sign-code shift: 7*l
mul.lo.u32 %r10, %r7, 7;
mov.u32 %r8, 0; // blk
$L_x2_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_x2_blk_end;
mul.wide.u32 %rd4, %r8, 66;
add.u64 %rd4, %rd3, %rd4; // block base
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// sub-block base = block + 2 + 8*ib
cvt.u64.u32 %rd5, %r9;
add.u64 %rd5, %rd4, %rd5;
// codebook index = byte l of aux0
cvt.u64.u32 %rd6, %r7;
add.u64 %rd6, %rd5, %rd6;
ld.global.u8 %r11, [%rd6];
// aux1 = u16[+4] | u16[+6] << 16 (2-byte aligned only)
ld.global.u16 %r12, [%rd5+4];
ld.global.u16 %r13, [%rd5+6];
shl.b32 %r13, %r13, 16;
or.b32 %r12, %r12, %r13;
// db = d * (0.5 + (aux1 >> 28)) * 0.25
shr.u32 %r14, %r12, 28;
cvt.rn.f32.u32 %f2, %r14;
add.f32 %f2, %f2, 0f3F000000; // +0.5
mul.f32 %f2, %f2, 0f3E800000; // *0.25 (exact: a power of two)
mul.f32 %f3, %f1, %f2; // db
// sign byte = ksigns[(aux1 >> 7l) & 127]
shr.u32 %r15, %r12, %r10;
and.b32 %r15, %r15, 127;
cvt.u64.u32 %rd7, %r15;
add.u64 %rd7, %rd15, %rd7;
ld.global.u8 %r25, [%rd7];
// grid entry: lo word = magnitudes 0..3, hi word = magnitudes 4..7
mul.wide.u32 %rd8, %r11, 8;
add.u64 %rd8, %rd10, %rd8;
ld.global.u32 %r22, [%rd8];
ld.global.u32 %r23, [%rd8+4];
// col0 = blk*256 + 32*ib + 8*l
shl.b32 %r26, %r8, 8;
shl.b32 %r27, %r6, 5;
add.u32 %r26, %r26, %r27;
shl.b32 %r28, %r7, 3;
add.u32 %r26, %r26, %r28;
mov.u32 %r29, 0; // j = 0..4
$L_x2_j:
setp.ge.u32 %p2, %r29, 4;
@%p2 bra $L_x2_j_end;
shl.b32 %r30, %r29, 3; // 8*j
shr.u32 %r31, %r22, %r30;
and.b32 %r31, %r31, 255; // m1 = magnitude j
shr.u32 %r32, %r23, %r30;
and.b32 %r32, %r32, 255; // m2 = magnitude j+4
// s1 = bit j of the sign byte, s2 = bit j+4
shr.u32 %r33, %r25, %r29;
and.b32 %r33, %r33, 1;
add.u32 %r34, %r29, 4;
shr.u32 %r35, %r25, %r34;
and.b32 %r35, %r35, 1;
cvt.rn.f32.u32 %f4, %r31;
mul.f32 %f4, %f4, %f3;
setp.ne.u32 %p3, %r33, 0;
@%p3 neg.f32 %f4, %f4;
cvt.rn.f32.u32 %f5, %r32;
mul.f32 %f5, %f5, %f3;
setp.ne.u32 %p4, %r35, 0;
@%p4 neg.f32 %f5, %f5;
// x[col0 + j]
add.u32 %r36, %r26, %r29;
setp.ge.u32 %p5, %r36, %r3;
@%p5 bra $L_x2_skip1;
mul.wide.u32 %rd12, %r36, 4;
add.u64 %rd12, %rd2, %rd12;
ld.global.f32 %f6, [%rd12];
fma.rn.f32 %f0, %f4, %f6, %f0;
$L_x2_skip1:
// x[col0 + j + 4]
add.u32 %r37, %r36, 4;
setp.ge.u32 %p6, %r37, %r3;
@%p6 bra $L_x2_skip2;
mul.wide.u32 %rd13, %r37, 4;
add.u64 %rd13, %rd2, %rd13;
ld.global.f32 %f7, [%rd13];
fma.rn.f32 %f0, %f5, %f7, %f0;
$L_x2_skip2:
add.u32 %r29, %r29, 1;
bra $L_x2_j;
$L_x2_j_end:
add.u32 %r8, %r8, 1;
bra $L_x2_blk;
$L_x2_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p7, %r0, 0;
@%p7 bra $L_x2_exit;
mul.wide.u32 %rd14, %r1, 4;
add.u64 %rd14, %rd0, %rd14;
st.global.f32 [%rd14], %f0;
$L_x2_exit:
ret;
}
",
);
ptx
}
fn generate_iq2_s_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
let grid = &crate::quantize::iq_grids::IQ2S_GRID;
let mut table = String::with_capacity(grid.len() * 24);
for (i, e) in grid.iter().enumerate() {
if i % 4 == 0 {
table.push_str("\n ");
}
table.push_str(&format!("{}, {}", *e as u32, (*e >> 32) as u32));
if i + 1 != grid.len() {
table.push_str(", ");
}
}
let mut ptx = String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// IQ2S_GRID: generated from quantize::iq_grids::IQ2S_GRID, 1024 u64 entries as
// 2048 little-endian u32 words (lo, hi).
.global .align 8 .u32 iq2s_grid_g[2048] = {",
);
ptx.push_str(&table);
ptx.push_str(
r"
};
.visible .entry iq2_s_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<48>;
.reg .u64 %rd<32>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_2s_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 82
mul.lo.u32 %r5, %r4, 82;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// ib = tid >> 2 (0..8), l = tid & 3 (0..4)
shr.u32 %r6, %r0, 2;
and.b32 %r7, %r0, 3;
mov.u64 %rd10, iq2s_grid_g;
mov.u32 %r8, 0; // blk
$L_2s_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_2s_blk_end;
mul.wide.u32 %rd4, %r8, 82;
add.u64 %rd4, %rd3, %rd4;
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// sc = scales[ib] (scales at +74)
cvt.u64.u32 %rd5, %r6;
add.u64 %rd5, %rd4, %rd5;
add.u64 %rd5, %rd5, 74;
ld.global.u8 %r10, [%rd5];
// nibble: l >> 1 selects it -- low for l = 0,1, high for l = 2,3
shr.u32 %r11, %r7, 1;
shl.b32 %r12, %r11, 2; // 0 or 4
shr.u32 %r13, %r10, %r12;
and.b32 %r13, %r13, 15;
// db = (d * (0.5 + nib)) * 0.25, in the reference's order
cvt.rn.f32.u32 %f2, %r13;
add.f32 %f2, %f2, 0f3F000000; // + 0.5
mul.f32 %f3, %f1, %f2; // d * (0.5 + nib)
mul.f32 %f3, %f3, 0f3E800000; // * 0.25 -> db
// qh_byte = qh[ib] (qh at +66)
cvt.u64.u32 %rd6, %r6;
add.u64 %rd6, %rd4, %rd6;
add.u64 %rd6, %rd6, 66;
ld.global.u8 %r14, [%rd6];
// lane = 4*ib + l
shl.b32 %r15, %r6, 2;
add.u32 %r15, %r15, %r7;
cvt.u64.u32 %rd7, %r15;
// idx low 8 bits = qs[4ib + l] (qs at +2)
add.u64 %rd8, %rd4, %rd7;
add.u64 %rd8, %rd8, 2;
ld.global.u8 %r17, [%rd8];
// idx high 2 bits = (qh_byte >> 2l) & 3, placed at bit 8
shl.b32 %r16, %r7, 1; // 2l
shr.u32 %r19, %r14, %r16;
and.b32 %r19, %r19, 3;
shl.b32 %r19, %r19, 8;
or.b32 %r17, %r17, %r19; // idx, 0..1023
// grid entry = 8 bytes: lo word bytes 0..3, hi word bytes 4..7
mul.wide.u32 %rd9, %r17, 8;
add.u64 %rd9, %rd10, %rd9;
ld.global.u32 %r22, [%rd9]; // g_lo
ld.global.u32 %r23, [%rd9+4]; // g_hi
// sign_byte = signs[4ib + l] (signs at +34)
add.u64 %rd11, %rd4, %rd7;
add.u64 %rd11, %rd11, 34;
ld.global.u8 %r25, [%rd11];
// col0 = blk*256 + 32*ib + 8*l
shl.b32 %r26, %r8, 8;
shl.b32 %r27, %r6, 5;
add.u32 %r26, %r26, %r27;
shl.b32 %r28, %r7, 3;
add.u32 %r26, %r26, %r28;
mov.u32 %r29, 0; // j = 0..4
$L_2s_j:
setp.ge.u32 %p2, %r29, 4;
@%p2 bra $L_2s_j_end;
shl.b32 %r30, %r29, 3; // 8*j
shr.u32 %r31, %r22, %r30;
and.b32 %r31, %r31, 255; // m_lo = byte j
shr.u32 %r32, %r23, %r30;
and.b32 %r32, %r32, 255; // m_hi = byte j+4
// sign bits: j for the lo byte, j+4 for the hi byte
shr.u32 %r33, %r25, %r29;
and.b32 %r33, %r33, 1;
add.u32 %r34, %r29, 4;
shr.u32 %r35, %r25, %r34;
and.b32 %r35, %r35, 1;
cvt.rn.f32.u32 %f4, %r31;
mul.f32 %f4, %f3, %f4; // db * m_lo
setp.ne.u32 %p3, %r33, 0;
@%p3 neg.f32 %f4, %f4;
cvt.rn.f32.u32 %f5, %r32;
mul.f32 %f5, %f3, %f5; // db * m_hi
setp.ne.u32 %p4, %r35, 0;
@%p4 neg.f32 %f5, %f5;
// x[col0 + j]
add.u32 %r36, %r26, %r29;
setp.ge.u32 %p5, %r36, %r3;
@%p5 bra $L_2s_skip1;
mul.wide.u32 %rd12, %r36, 4;
add.u64 %rd12, %rd2, %rd12;
ld.global.f32 %f6, [%rd12];
fma.rn.f32 %f0, %f4, %f6, %f0;
$L_2s_skip1:
// x[col0 + j + 4]
add.u32 %r37, %r36, 4;
setp.ge.u32 %p6, %r37, %r3;
@%p6 bra $L_2s_skip2;
mul.wide.u32 %rd13, %r37, 4;
add.u64 %rd13, %rd2, %rd13;
ld.global.f32 %f7, [%rd13];
fma.rn.f32 %f0, %f5, %f7, %f0;
$L_2s_skip2:
add.u32 %r29, %r29, 1;
bra $L_2s_j;
$L_2s_j_end:
add.u32 %r8, %r8, 1;
bra $L_2s_blk;
$L_2s_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p7, %r0, 0;
@%p7 bra $L_2s_exit;
mul.wide.u32 %rd14, %r1, 4;
add.u64 %rd14, %rd0, %rd14;
st.global.f32 [%rd14], %f0;
$L_2s_exit:
ret;
}
",
);
ptx
}
fn generate_iq3_xxs_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
let grid = crate::quantize::iq_grids::IQ3XXS_GRID
.iter()
.map(u32::to_string)
.collect::<Vec<_>>()
.join(", ");
let signs = crate::quantize::iq_grids::KSIGNS_IQ2XS
.iter()
.map(u8::to_string)
.collect::<Vec<_>>()
.join(", ");
let mut ptx = String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// IQ3XXS_GRID: 256 packed 4-magnitude entries. GENERATED from
// quantize::iq_grids::IQ3XXS_GRID when this PTX is built, never hand-copied.
.global .align 4 .u32 iq3xxs_grid_g[256] = {",
);
ptx.push_str(&grid);
ptx.push_str(
r"};
// KSIGNS_IQ2XS: 128 sign bytes. GENERATED from quantize::iq_grids::KSIGNS_IQ2XS.
.global .align 1 .u8 iq3xxs_ksigns_g[128] = {",
);
ptx.push_str(&signs);
ptx.push_str(
r"};
.visible .entry iq3_xxs_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<48>;
.reg .u64 %rd<32>;
.reg .f32 %f<24>;
.reg .b16 %h<4>;
.reg .pred %p<12>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_x3_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 256)
add.u32 %r4, %r3, 255;
shr.u32 %r4, %r4, 8;
// row_base = w_ptr + ctaid * nb * 98
mul.lo.u32 %r5, %r4, 98;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// ib = tid >> 2 (0..8), l = tid & 3 (0..4)
shr.u32 %r6, %r0, 2;
and.b32 %r7, %r0, 3;
mov.u64 %rd10, iq3xxs_grid_g;
mov.u64 %rd15, iq3xxs_ksigns_g;
// this lane's scale/sign word offset: 66 + 4*ib
shl.b32 %r9, %r6, 2;
add.u32 %r9, %r9, 66;
// this lane's grid-index pair offset: 2 + 8*ib + 2*l
shl.b32 %r16, %r6, 3;
shl.b32 %r19, %r7, 1;
add.u32 %r16, %r16, %r19;
add.u32 %r16, %r16, 2;
// this lane's sign-code shift: 7*l
mul.lo.u32 %r10, %r7, 7;
mov.u32 %r8, 0; // blk
$L_x3_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_x3_blk_end;
mul.wide.u32 %rd4, %r8, 98;
add.u64 %rd4, %rd3, %rd4; // block base
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// aux = u16[+66+4ib] | u16[+68+4ib] << 16 (2-byte aligned only)
cvt.u64.u32 %rd5, %r9;
add.u64 %rd5, %rd4, %rd5;
ld.global.u16 %r12, [%rd5];
ld.global.u16 %r13, [%rd5+2];
shl.b32 %r13, %r13, 16;
or.b32 %r12, %r12, %r13;
// db = d * (0.5 + (aux >> 28)) * 0.5
shr.u32 %r14, %r12, 28;
cvt.rn.f32.u32 %f2, %r14;
add.f32 %f2, %f2, 0f3F000000; // +0.5
mul.f32 %f2, %f2, 0f3F000000; // *0.5 (exact: a power of two)
mul.f32 %f3, %f1, %f2; // db
// sign byte = ksigns[(aux >> 7l) & 127]
shr.u32 %r15, %r12, %r10;
and.b32 %r15, %r15, 127;
cvt.u64.u32 %rd7, %r15;
add.u64 %rd7, %rd15, %rd7;
ld.global.u8 %r25, [%rd7];
// grid indices qs[8ib+2l], qs[8ib+2l+1]
cvt.u64.u32 %rd6, %r16;
add.u64 %rd6, %rd4, %rd6;
ld.global.u8 %r17, [%rd6];
ld.global.u8 %r18, [%rd6+1];
// g1 = grid[i1] (columns j), g2 = grid[i2] (columns j+4)
mul.wide.u32 %rd8, %r17, 4;
add.u64 %rd8, %rd10, %rd8;
ld.global.u32 %r22, [%rd8];
mul.wide.u32 %rd9, %r18, 4;
add.u64 %rd9, %rd10, %rd9;
ld.global.u32 %r23, [%rd9];
// col0 = blk*256 + 32*ib + 8*l
shl.b32 %r26, %r8, 8;
shl.b32 %r27, %r6, 5;
add.u32 %r26, %r26, %r27;
shl.b32 %r28, %r7, 3;
add.u32 %r26, %r26, %r28;
mov.u32 %r29, 0; // j = 0..4
$L_x3_j:
setp.ge.u32 %p2, %r29, 4;
@%p2 bra $L_x3_j_end;
shl.b32 %r30, %r29, 3; // 8*j
shr.u32 %r31, %r22, %r30;
and.b32 %r31, %r31, 255; // m1
shr.u32 %r32, %r23, %r30;
and.b32 %r32, %r32, 255; // m2
// s1 = bit j of the sign byte, s2 = bit j+4
shr.u32 %r33, %r25, %r29;
and.b32 %r33, %r33, 1;
add.u32 %r34, %r29, 4;
shr.u32 %r35, %r25, %r34;
and.b32 %r35, %r35, 1;
cvt.rn.f32.u32 %f4, %r31;
mul.f32 %f4, %f4, %f3;
setp.ne.u32 %p3, %r33, 0;
@%p3 neg.f32 %f4, %f4;
cvt.rn.f32.u32 %f5, %r32;
mul.f32 %f5, %f5, %f3;
setp.ne.u32 %p4, %r35, 0;
@%p4 neg.f32 %f5, %f5;
// x[col0 + j]
add.u32 %r36, %r26, %r29;
setp.ge.u32 %p5, %r36, %r3;
@%p5 bra $L_x3_skip1;
mul.wide.u32 %rd12, %r36, 4;
add.u64 %rd12, %rd2, %rd12;
ld.global.f32 %f6, [%rd12];
fma.rn.f32 %f0, %f4, %f6, %f0;
$L_x3_skip1:
// x[col0 + j + 4]
add.u32 %r37, %r36, 4;
setp.ge.u32 %p6, %r37, %r3;
@%p6 bra $L_x3_skip2;
mul.wide.u32 %rd13, %r37, 4;
add.u64 %rd13, %rd2, %rd13;
ld.global.f32 %f7, [%rd13];
fma.rn.f32 %f0, %f5, %f7, %f0;
$L_x3_skip2:
add.u32 %r29, %r29, 1;
bra $L_x3_j;
$L_x3_j_end:
add.u32 %r8, %r8, 1;
bra $L_x3_blk;
$L_x3_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p7, %r0, 0;
@%p7 bra $L_x3_exit;
mul.wide.u32 %rd14, %r1, 4;
add.u64 %rd14, %rd0, %rd14;
st.global.f32 [%rd14], %f0;
$L_x3_exit:
ret;
}
",
);
ptx
}
fn generate_iq4_nl_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// The 16 non-linear IQ4_NL levels: quantize::iq_grids::KVALUES_IQ4NL.
.global .align 4 .s32 kvalues_iq4nl_b[16] = {-127, -104, -83, -65, -49, -35, -22, -10, 1, 13, 25, 38, 53, 69, 89, 113};
.visible .entry iq4_nl_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<40>;
.reg .u64 %rd<24>;
.reg .f32 %f<20>;
.reg .b16 %h<4>;
.reg .pred %p<10>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_nl_exit;
mov.f32 %f0, 0f00000000;
// nb = ceil(k_dim / 32)
add.u32 %r4, %r3, 31;
shr.u32 %r4, %r4, 5;
// row_base = w_ptr + ctaid * nb * 18
mul.lo.u32 %r5, %r4, 18;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// per-thread: jlow = tid & 15 (byte index), jhalf = tid >> 4 (nibble select)
and.b32 %r7, %r0, 15;
shr.u32 %r6, %r0, 4;
// nibble shift = jhalf * 4
shl.b32 %r19, %r6, 2;
// codebook base
mov.u64 %rd10, kvalues_iq4nl_b;
mov.u32 %r8, 0;
$L_nl_blk:
setp.ge.u32 %p1, %r8, %r4;
@%p1 bra $L_nl_blk_end;
// blk_addr = row_base + blk * 18
mul.wide.u32 %rd4, %r8, 18;
add.u64 %rd4, %rd3, %rd4;
// d (f16 at +0)
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// byte = qs[jlow], qs starts at +2
cvt.u64.u32 %rd7, %r7;
add.u64 %rd7, %rd4, %rd7;
add.u64 %rd7, %rd7, 2;
ld.global.u8 %r18, [%rd7];
// nib = jhalf ? (byte >> 4) : (byte & 0xf)
shr.u32 %r20, %r18, %r19;
and.b32 %r20, %r20, 15;
// w = d * kvalues_iq4nl[nib]
mul.wide.u32 %rd8, %r20, 4;
add.u64 %rd8, %rd10, %rd8;
ld.global.s32 %r21, [%rd8];
cvt.rn.f32.s32 %f4, %r21;
mul.f32 %f5, %f1, %f4;
// x_idx = blk*32 + tid
shl.b32 %r22, %r8, 5;
add.u32 %r22, %r22, %r0;
setp.ge.u32 %p3, %r22, %r3;
@%p3 bra $L_nl_skip;
mul.wide.u32 %rd9, %r22, 4;
add.u64 %rd9, %rd2, %rd9;
ld.global.f32 %f6, [%rd9];
fma.rn.f32 %f0, %f5, %f6, %f0;
$L_nl_skip:
add.u32 %r8, %r8, 1;
bra $L_nl_blk;
$L_nl_blk_end:
shfl.sync.down.b32 %f10, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f10;
shfl.sync.down.b32 %f11, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f11;
shfl.sync.down.b32 %f12, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f12;
shfl.sync.down.b32 %f13, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f13;
shfl.sync.down.b32 %f14, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f14;
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_nl_exit;
mul.wide.u32 %rd11, %r1, 4;
add.u64 %rd11, %rd0, %rd11;
st.global.f32 [%rd11], %f0;
$L_nl_exit:
ret;
}
",
)
}
fn generate_f16_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// F16 GEMV, row-major: y[row] = sum_i f16_to_f32(w[row*k + i]) * x[i]
.visible .entry f16_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<20>;
.reg .u64 %rd<16>;
.reg .f32 %f<16>;
.reg .b16 %h<4>;
.reg .pred %p<8>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
// Rows beyond n_dim do no work (grid is padded to whole warps).
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_exit;
mov.f32 %f0, 0f00000000;
// rd3 = row_base = w_ptr + ctaid * k_dim * 2 (2 bytes per f16)
shl.b32 %r4, %r3, 1;
mul.wide.u32 %rd3, %r1, %r4;
add.u64 %rd3, %rd1, %rd3;
// i = tid, stride 32
mov.u32 %r5, %r0;
$L_loop:
setp.ge.u32 %p1, %r5, %r3;
@%p1 bra $L_loop_end;
// w = f16_to_f32(row_base[i])
mul.wide.u32 %rd4, %r5, 2;
add.u64 %rd4, %rd3, %rd4;
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// x = x_ptr[i]
mul.wide.u32 %rd5, %r5, 4;
add.u64 %rd5, %rd2, %rd5;
ld.global.f32 %f2, [%rd5];
fma.rn.f32 %f0, %f1, %f2, %f0;
add.u32 %r5, %r5, 32;
bra $L_loop;
$L_loop_end:
// Warp reduction: identical idiom to the block-quantized GEMVs here.
shfl.sync.down.b32 %f4, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f4;
shfl.sync.down.b32 %f5, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f5;
shfl.sync.down.b32 %f6, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f6;
shfl.sync.down.b32 %f7, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f7;
shfl.sync.down.b32 %f8, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f8;
setp.ne.u32 %p2, %r0, 0;
@%p2 bra $L_exit;
mul.wide.u32 %rd6, %r1, 4;
add.u64 %rd6, %rd0, %rd6;
st.global.f32 [%rd6], %f0;
$L_exit:
ret;
}
",
)
}
fn generate_bf16_gemv_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// BF16 GEMV, row-major: y[row] = sum_i bf16_to_f32(w[row*k + i]) * x[i]
// bf16_to_f32(bits) == f32::from_bits(bits << 16), exact.
.visible .entry bf16_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<20>;
.reg .u64 %rd<16>;
.reg .f32 %f<16>;
.reg .pred %p<8>;
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
// Rows beyond n_dim do no work (grid is padded to whole warps).
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_exit;
mov.f32 %f0, 0f00000000;
// rd3 = row_base = w_ptr + ctaid * k_dim * 2 (2 bytes per bf16)
shl.b32 %r4, %r3, 1;
mul.wide.u32 %rd3, %r1, %r4;
add.u64 %rd3, %rd1, %rd3;
// i = tid, stride 32
mov.u32 %r5, %r0;
$L_loop:
setp.ge.u32 %p1, %r5, %r3;
@%p1 bra $L_loop_end;
// w = bf16_to_f32(row_base[i]): zero-extending 16-bit load, shift into the
// high half, reinterpret. mov.b32 between .u32 and .f32 is a bitcast.
mul.wide.u32 %rd4, %r5, 2;
add.u64 %rd4, %rd3, %rd4;
ld.global.u16 %r6, [%rd4];
shl.b32 %r7, %r6, 16;
mov.b32 %f1, %r7;
// x = x_ptr[i]
mul.wide.u32 %rd5, %r5, 4;
add.u64 %rd5, %rd2, %rd5;
ld.global.f32 %f2, [%rd5];
fma.rn.f32 %f0, %f1, %f2, %f0;
add.u32 %r5, %r5, 32;
bra $L_loop;
$L_loop_end:
// Warp reduction: identical idiom to the block-quantized GEMVs here.
shfl.sync.down.b32 %f4, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f4;
shfl.sync.down.b32 %f5, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f5;
shfl.sync.down.b32 %f6, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f6;
shfl.sync.down.b32 %f7, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f7;
shfl.sync.down.b32 %f8, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f8;
setp.ne.u32 %p2, %r0, 0;
@%p2 bra $L_exit;
mul.wide.u32 %rd6, %r1, 4;
add.u64 %rd6, %rd0, %rd6;
st.global.f32 [%rd6], %f0;
$L_exit:
ret;
}
",
)
}
fn generate_q5_0_candle_ptx(k: u32, n: u32) -> String {
let _ = (k, n);
String::from(
r"
.version 7.5
.target sm_70
.address_size 64
// BUG-GGUF-002 FIX: Q5_0 GEMV with candle nibble layout
// Each warp (32 threads) computes one output element
// Thread 0-15: use low nibbles from bytes 0-15, qh bits 0-15
// Thread 16-31: use high nibbles from bytes 0-15, qh bits 16-31
.visible .entry q5_0_gemv_warp_reduce(
.param .u64 y_ptr,
.param .u64 w_ptr,
.param .u64 x_ptr,
.param .u32 k_dim,
.param .u32 n_dim
)
{
.reg .u32 %r<40>;
.reg .u64 %rd<20>;
.reg .f32 %f<16>;
.reg .b16 %h<4>;
.reg .pred %p<8>;
// r0=tid, r1=ctaid, r2=n_dim, r3=k_dim
mov.u32 %r0, %tid.x;
mov.u32 %r1, %ctaid.x;
ld.param.u32 %r2, [n_dim];
ld.param.u32 %r3, [k_dim];
ld.param.u64 %rd0, [y_ptr];
ld.param.u64 %rd1, [w_ptr];
ld.param.u64 %rd2, [x_ptr];
// Bounds check: if ctaid >= n_dim, exit
setp.ge.u32 %p0, %r1, %r2;
@%p0 bra $L_exit;
// f0 = accumulator
mov.f32 %f0, 0f00000000;
// r4 = num_blocks = ceil(k_dim / 32)
add.u32 %r4, %r3, 31;
shr.u32 %r4, %r4, 5;
// rd3 = row_base = w_ptr + ctaid * num_blocks * 22
mul.lo.u32 %r5, %r4, 22;
mul.wide.u32 %rd3, %r1, %r5;
add.u64 %rd3, %rd1, %rd3;
// r6 = blk_idx (loop counter)
mov.u32 %r6, 0;
$L_blk_loop:
setp.ge.u32 %p1, %r6, %r4;
@%p1 bra $L_blk_loop_end;
// rd4 = blk_addr = row_base + blk_idx * 22
mul.wide.u32 %rd4, %r6, 22;
add.u64 %rd4, %rd3, %rd4;
// f1 = scale d (fp16 at offset 0) - use b16 register for f16 conversion
ld.global.b16 %h0, [%rd4];
cvt.f32.f16 %f1, %h0;
// Load qh (4 bytes at offset 2) using byte loads for unaligned access
add.u64 %rd5, %rd4, 2;
ld.global.u8 %r20, [%rd5];
add.u64 %rd6, %rd4, 3;
ld.global.u8 %r21, [%rd6];
add.u64 %rd7, %rd4, 4;
ld.global.u8 %r22, [%rd7];
add.u64 %rd8, %rd4, 5;
ld.global.u8 %r23, [%rd8];
// Combine: qh = r20 | (r21 << 8) | (r22 << 16) | (r23 << 24)
shl.b32 %r24, %r21, 8;
shl.b32 %r25, %r22, 16;
shl.b32 %r26, %r23, 24;
or.b32 %r27, %r20, %r24;
or.b32 %r28, %r27, %r25;
or.b32 %r8, %r28, %r26; // r8 = qh
// rd9 = qs_base = blk_addr + 6
add.u64 %rd9, %rd4, 6;
// CANDLE LAYOUT:
// Thread 0-15 read bytes 0-15 (low nibbles -> positions 0-15), qh bits 0-15
// Thread 16-31 read bytes 0-15 (high nibbles -> positions 16-31), qh bits 16-31
// r9 = byte_idx = tid < 16 ? tid : tid - 16
setp.ge.u32 %p2, %r0, 16;
mov.u32 %r9, %r0;
@%p2 sub.u32 %r9, %r0, 16;
// Load byte from qs[byte_idx]
cvt.u64.u32 %rd10, %r9;
add.u64 %rd10, %rd9, %rd10;
ld.global.u8 %r10, [%rd10];
// r11 = nibble value
// Threads 0-15: low nibble (byte & 0xF)
// Threads 16-31: high nibble (byte >> 4)
mov.u32 %r11, %r10;
@%p2 shr.u32 %r11, %r10, 4;
and.b32 %r11, %r11, 15;
// Extract high bit: (qh >> tid) & 1
// For candle layout, threads 0-15 use qh bits 0-15, threads 16-31 use qh bits 16-31
shr.b32 %r12, %r8, %r0;
and.b32 %r12, %r12, 1;
// Combine: q5 = nibble | (high_bit << 4)
shl.b32 %r13, %r12, 4;
or.b32 %r14, %r11, %r13;
// r15 = centered value = q5 - 16 (as signed)
sub.u32 %r15, %r14, 16;
// f2 = dequantized = d * centered
cvt.rn.f32.s32 %f2, %r15;
mul.f32 %f2, %f1, %f2;
// r16 = x_idx = blk_idx * 32 + tid
shl.b32 %r16, %r6, 5;
add.u32 %r16, %r16, %r0;
// Bounds check for last block
setp.ge.u32 %p3, %r16, %r3;
@%p3 bra $L_skip_mul;
// f3 = x[x_idx]
cvt.u64.u32 %rd11, %r16;
shl.b64 %rd11, %rd11, 2;
add.u64 %rd11, %rd2, %rd11;
ld.global.f32 %f3, [%rd11];
// f0 += f2 * f3
fma.rn.f32 %f0, %f2, %f3, %f0;
$L_skip_mul:
add.u32 %r6, %r6, 1;
bra $L_blk_loop;
$L_blk_loop_end:
// Warp reduction using shfl.sync.down
shfl.sync.down.b32 %f4, %f0, 16, 31, 0xffffffff;
add.f32 %f0, %f0, %f4;
shfl.sync.down.b32 %f5, %f0, 8, 31, 0xffffffff;
add.f32 %f0, %f0, %f5;
shfl.sync.down.b32 %f6, %f0, 4, 31, 0xffffffff;
add.f32 %f0, %f0, %f6;
shfl.sync.down.b32 %f7, %f0, 2, 31, 0xffffffff;
add.f32 %f0, %f0, %f7;
shfl.sync.down.b32 %f8, %f0, 1, 31, 0xffffffff;
add.f32 %f0, %f0, %f8;
// Thread 0 writes result
setp.ne.u32 %p4, %r0, 0;
@%p4 bra $L_exit;
// y[ctaid] = f0
mul.wide.u32 %rd12, %r1, 4;
add.u64 %rd12, %rd0, %rd12;
st.global.f32 [%rd12], %f0;
$L_exit:
ret;
}
",
)
}