| // | |
| // 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; | |
| } | |