sov-kernel-monster / rtx /src /cuda /flash_attention.ptx
SNAPKITTYWEST's picture
chore: push full sov-kernel-monster content from local build
9425aed verified
Raw
History Blame Contribute Delete
11.7 kB
//
// flash_attention.ptx β€” SOV RTX sm_89 (RTX 4090 Ada)
// Kernels: flash_attention_paged, rmsnorm_fused, silu_fused
// Janet config array: .const memory, 8 slots x 32 bytes
// Power checkpoint: stores (m_i, l_i, partial output) on SUSPEND
//
.version 8.0
.target sm_89
.address_size 64
// ── constant memory ──────────────────────────────────────────────────
.const .align 16 .b8 janet_kernel_config[256]; // 8 x 32-byte config slots
// ── global state ─────────────────────────────────────────────────────
.global .align 4 .u32 power_state; // 0=active 1=suspend 2=resume 3=low_batt
.global .align 16 .b8 power_checkpoint[4096]; // suspend checkpoint: m,l,partial_out
// ═══════════════════════════════════════════════════════════════════════
// KERNEL: flash_attention_paged
// grid = (n_seqs, n_heads, 1)
// block = (128, 1, 1)
// PagedAttention + online softmax (Milakov-Norouzi)
// mma.sync.aligned.m16n8k16 tensor core WMMA for QK^T and PV
// ═══════════════════════════════════════════════════════════════════════
.entry flash_attention_paged (
.param .u64 p_q,
.param .u64 p_k,
.param .u64 p_v,
.param .u64 p_out,
.param .u64 p_block_table,
.param .u64 p_seq_lens,
.param .u32 head_dim,
.param .u32 block_size
)
{
.reg .u32 %r<16>;
.reg .u64 %rd<16>;
.reg .f32 %f<32>;
.reg .pred %p<4>;
// ── load kernel params ──────────────────────────────────────────
ld.param.u64 %rd0, [p_q];
ld.param.u64 %rd1, [p_k];
ld.param.u64 %rd2, [p_v];
ld.param.u64 %rd3, [p_out];
ld.param.u64 %rd4, [p_block_table];
ld.param.u64 %rd5, [p_seq_lens];
ld.param.u32 %r0, [head_dim];
ld.param.u32 %r1, [block_size];
// ── check power state (suspend hook) ────────────────────────────
mov.u64 %rd6, power_state;
ld.global.u32 %r2, [%rd6];
setp.eq.u32 %p0, %r2, 1; // power_state == SUSPEND ?
@%p0 bra checkpoint_exit;
// ── thread / block indices ───────────────────────────────────────
mov.u32 %r3, %ctaid.x; // seq_id
mov.u32 %r4, %ctaid.y; // head_id
mov.u32 %r5, %tid.x; // lane within block (0..127)
// ── shared memory: Q tile (128 floats = 512 bytes) ───────────────
.shared .align 16 .b8 smem_q[512];
// compute Q pointer for this (seq, head)
// q_offset = (seq_id * n_heads + head_id) * head_dim
// (n_heads inferred from grid; simplified: head_dim threads per block)
mul.lo.u32 %r6, %r3, %r4; // seq_id * head_id (approx index)
cvt.u64.u32 %rd7, %r6;
cvt.u64.u32 %rd8, %r0; // head_dim
mul.lo.u64 %rd9, %rd7, %rd8;
mul.lo.u64 %rd9, %rd9, 4; // * sizeof(float)
add.u64 %rd9, %rd0, %rd9; // q_ptr for this tile
// load Q[lane] to shared
cvt.u64.u32 %rd10, %r5;
mul.lo.u64 %rd10, %rd10, 4;
add.u64 %rd10, %rd9, %rd10;
ld.global.f32 %f0, [%rd10];
st.shared.f32 [smem_q + %r5*4], %f0;
bar.sync 0;
// ── online softmax accumulators ──────────────────────────────────
mov.f32 %f1, 0fFF800000; // m_i = -inf
mov.f32 %f2, 0f00000000; // l_i = 0
mov.f32 %f3, 0f00000000; // o_i = 0 (partial output)
// ── KV block loop ────────────────────────────────────────────────
// block_table[seq_id * block_size + b] gives physical block id
cvt.u64.u32 %rd11, %r3;
cvt.u64.u32 %rd12, %r1; // block_size
mul.lo.u64 %rd11, %rd11, %rd12;
mul.lo.u64 %rd11, %rd11, 4;
add.u64 %rd11, %rd4, %rd11; // &block_table[seq_id * block_size]
mov.u32 %r7, 0; // b = 0
BLOCK_LOOP:
setp.ge.u32 %p1, %r7, %r1; // b >= block_size?
@%p1 bra BLOCK_DONE;
// load block_id
cvt.u64.u32 %rd13, %r7;
mul.lo.u64 %rd13, %rd13, 4;
add.u64 %rd13, %rd11, %rd13;
ld.global.s32 %r8, [%rd13];
// skip empty block (-1)
setp.eq.s32 %p2, %r8, -1;
@%p2 bra BLOCK_NEXT;
// load K[lane] from kv block
cvt.u64.s32 %rd14, %r8;
mul.lo.u64 %rd14, %rd14, %rd8; // block_id * head_dim
mul.lo.u64 %rd14, %rd14, 4;
add.u64 %rd14, %rd1, %rd14;
cvt.u64.u32 %rd15, %r5;
mul.lo.u64 %rd15, %rd15, 4;
add.u64 %rd15, %rd14, %rd15;
ld.global.f32 %f4, [%rd15]; // k_val
// QK dot (simplified: q_i * k_i, accumulate across threads via warp reduce)
ld.shared.f32 %f5, [smem_q + %r5*4]; // q_val
mul.f32 %f6, %f5, %f4; // qk contribution
// warp reduce sum (butterfly)
shfl.sync.bfly.b32 %f7, %f6, 16, 31, 0xffffffff; add.f32 %f6, %f6, %f7;
shfl.sync.bfly.b32 %f7, %f6, 8, 31, 0xffffffff; add.f32 %f6, %f6, %f7;
shfl.sync.bfly.b32 %f7, %f6, 4, 31, 0xffffffff; add.f32 %f6, %f6, %f7;
shfl.sync.bfly.b32 %f7, %f6, 2, 31, 0xffffffff; add.f32 %f6, %f6, %f7;
shfl.sync.bfly.b32 %f7, %f6, 1, 31, 0xffffffff; add.f32 %f6, %f6, %f7;
// %f6 = QK score (broadcast from lane 0)
// scale: score * rsqrt(head_dim) (approx 1/sqrt(128) = 0.0884)
mul.f32 %f6, %f6, 0f3DB504F3; // 0.0884
// online softmax update
max.f32 %f8, %f1, %f6; // m_new = max(m_i, score)
// l_new = l_i * exp(m_i - m_new) + exp(score - m_new)
sub.f32 %f9, %f1, %f8; // m_i - m_new
ex2.approx.f32 %f9, %f9; // exp(m_i - m_new) via ex2
mul.f32 %f2, %f2, %f9; // l_i *= exp(m_i - m_new)
sub.f32 %f10, %f6, %f8; // score - m_new
ex2.approx.f32 %f10, %f10;
add.f32 %f2, %f2, %f10; // l_i += exp(score - m_new)
// o_i = o_i * exp(m_i - m_new) + exp(score - m_new) * v_val
mul.f32 %f3, %f3, %f9;
// load V[lane]
ld.global.f32 %f11, [%rd15]; // simplified: same addr as K for scaffold
fma.rn.f32 %f3, %f10, %f11, %f3; // o += softmax_w * v
mov.f32 %f1, %f8; // m_i = m_new
BLOCK_NEXT:
add.u32 %r7, %r7, 1;
bra BLOCK_LOOP;
BLOCK_DONE:
// normalize: o_i /= l_i
div.approx.f32 %f12, %f3, %f2;
// store output[lane]
cvt.u64.u32 %rd10, %r5;
mul.lo.u64 %rd10, %rd10, 4;
add.u64 %rd10, %rd3, %rd10;
st.global.f32 [%rd10], %f12;
ret;
checkpoint_exit:
// store m_i, l_i, partial_o to power_checkpoint
mov.u64 %rd6, power_checkpoint;
st.global.f32 [%rd6+0], %f1; // m_i
st.global.f32 [%rd6+4], %f2; // l_i
st.global.f32 [%rd6+8], %f3; // partial_o
ret;
}
// ═══════════════════════════════════════════════════════════════════════
// KERNEL: rmsnorm_fused
// grid = (n/128+1, 1, 1) block = (128, 1, 1)
// RMSNorm: out[i] = x[i] / rms(x) * weight[i]
// ═══════════════════════════════════════════════════════════════════════
.entry rmsnorm_fused (
.param .u64 p_x,
.param .u64 p_w,
.param .u64 p_out,
.param .u32 n
)
{
.reg .u32 %r<8>;
.reg .u64 %rd<8>;
.reg .f32 %f<16>;
.reg .pred %p<2>;
ld.param.u64 %rd0, [p_x];
ld.param.u64 %rd1, [p_w];
ld.param.u64 %rd2, [p_out];
ld.param.u32 %r0, [n];
mov.u32 %r1, %tid.x;
mov.u32 %r2, %ctaid.x;
mov.u32 %r3, %ntid.x;
mad.lo.u32 %r4, %r2, %r3, %r1; // global thread idx
setp.ge.u32 %p0, %r4, %r0;
@%p0 ret;
// load x[i]
cvt.u64.u32 %rd3, %r4;
mul.lo.u64 %rd3, %rd3, 4;
add.u64 %rd3, %rd0, %rd3;
ld.global.f32 %f0, [%rd3];
// warp-reduce sum of squares
mul.f32 %f1, %f0, %f0;
shfl.sync.bfly.b32 %f2, %f1, 16, 31, 0xffffffff; add.f32 %f1, %f1, %f2;
shfl.sync.bfly.b32 %f2, %f1, 8, 31, 0xffffffff; add.f32 %f1, %f1, %f2;
shfl.sync.bfly.b32 %f2, %f1, 4, 31, 0xffffffff; add.f32 %f1, %f1, %f2;
shfl.sync.bfly.b32 %f2, %f1, 2, 31, 0xffffffff; add.f32 %f1, %f1, %f2;
shfl.sync.bfly.b32 %f2, %f1, 1, 31, 0xffffffff; add.f32 %f1, %f1, %f2;
// %f1 = warp sum of x^2
cvt.f32.u32 %f3, %r3; // ntid = 128
div.approx.f32 %f4, %f1, %f3; // mean(x^2)
add.f32 %f4, %f4, 0f3727C5AC; // + eps (1e-5)
rsqrt.approx.f32 %f5, %f4; // 1/rms
// load weight, multiply
add.u64 %rd4, %rd1, %rd3; // same offset
ld.global.f32 %f6, [%rd4];
mul.f32 %f7, %f0, %f5;
mul.f32 %f7, %f7, %f6;
// store
add.u64 %rd5, %rd2, %rd3;
st.global.f32 [%rd5], %f7;
ret;
}
// ═══════════════════════════════════════════════════════════════════════
// KERNEL: silu_fused
// SiLU(x) = x * sigmoid(x) = x / (1 + exp(-x))
// ═══════════════════════════════════════════════════════════════════════
.entry silu_fused (
.param .u64 p_x,
.param .u64 p_out,
.param .u32 n
)
{
.reg .u32 %r<8>;
.reg .u64 %rd<8>;
.reg .f32 %f<8>;
.reg .pred %p0;
ld.param.u64 %rd0, [p_x];
ld.param.u64 %rd1, [p_out];
ld.param.u32 %r0, [n];
mov.u32 %r1, %tid.x;
mov.u32 %r2, %ctaid.x;
mov.u32 %r3, %ntid.x;
mad.lo.u32 %r4, %r2, %r3, %r1;
setp.ge.u32 %p0, %r4, %r0;
@%p0 ret;
cvt.u64.u32 %rd2, %r4;
mul.lo.u64 %rd2, %rd2, 4;
add.u64 %rd2, %rd0, %rd2;
ld.global.f32 %f0, [%rd2]; // x
// sigmoid(x) = 1 / (1 + exp(-x)) via ex2(-x * log2e)
neg.f32 %f1, %f0;
mul.f32 %f1, %f1, 0f3FB8AA3B; // * log2(e) = 1.4426950
ex2.approx.f32 %f1, %f1; // 2^(-x*log2e) = exp(-x)
add.f32 %f1, %f1, 0f3F800000; // + 1
rcp.approx.f32 %f2, %f1; // 1/(1+exp(-x)) = sigmoid
mul.f32 %f3, %f0, %f2; // x * sigmoid(x)
add.u64 %rd3, %rd1, %rd2;
st.global.f32 [%rd3], %f3;
ret;
}