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