code wiki / (root) / nx_flash_attention_test.nx

nx_flash_attention_test.nx source

↩ module page · 121 lines · 5628 B

1// nx_flash_attention_test.nx -- tiled-attention equivalence + memory win. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_tensor.nx" 6import "nx_attention.nx" 7import "nx_flash_attention.nx" 8import "nx_numeric_oracle.nx" 9 10func main() -> nx_int { 11 let err: *i64 = (sys_mmap(8)) as *i64 12 13 // ===== 2 queries x 4 KV positions x head_dim 3 ============== 14 // 15 // Same shapes as nx_attention_test so we have a working 16 // reference oracle. Flash output must match standard attention 17 // within epsilon (Q10 rounding can drift slightly due to 18 // online-softmax accumulation order). 19 20 let sh_q: *i64 = (sys_mmap(16)) as *i64 21 sh_q[0] = 2; sh_q[1] = 3 22 let q: *NxTensor = nx_t_alloc(NX_DT_I64, sh_q, 2, err) 23 nx_t_fill_zero(q) 24 let pq: *i64 = q.storage as *i64 25 pq[0] = 1; pq[1] = 0; pq[2] = 0 // Q[0] = [1,0,0] 26 pq[3] = 0; pq[4] = 1; pq[5] = 0 // Q[1] = [0,1,0] 27 28 let sh_kv: *i64 = (sys_mmap(16)) as *i64 29 sh_kv[0] = 4; sh_kv[1] = 3 30 let k: *NxTensor = nx_t_alloc(NX_DT_I64, sh_kv, 2, err) 31 nx_t_fill_zero(k) 32 let pk: *i64 = k.storage as *i64 33 pk[0] = 2; pk[1] = 0; pk[2] = 0 // K[0] = [2,0,0] Q[0].K[0]=2 34 pk[3] = 0; pk[4] = 3; pk[5] = 0 // K[1] = [0,3,0] Q[1].K[1]=3 35 pk[6] = 1; pk[7] = 1; pk[8] = 0 // K[2] = [1,1,0] both Q[0],Q[1] dot=1 36 pk[9] = 0; pk[10] = 0; pk[11] = 1 // K[3] = [0,0,1] neither dot=0 37 38 let v: *NxTensor = nx_t_alloc(NX_DT_I64, sh_kv, 2, err) 39 let pv: *i64 = v.storage as *i64 40 pv[0] = 10; pv[1] = 20; pv[2] = 30 // V[0] 41 pv[3] = 40; pv[4] = 50; pv[5] = 60 // V[1] 42 pv[6] = 70; pv[7] = 80; pv[8] = 90 // V[2] 43 pv[9] = 100; pv[10] = 110; pv[11] = 120 // V[3] 44 45 let sh_out: *i64 = (sys_mmap(16)) as *i64 46 sh_out[0] = 2; sh_out[1] = 3 47 let out_std: *NxTensor = nx_t_alloc(NX_DT_I64, sh_out, 2, err) 48 let out_flash: *NxTensor = nx_t_alloc(NX_DT_I64, sh_out, 2, err) 49 50 // ===== Reference: standard attention forward ================ 51 let scale_q10: nx_int = 512 // (Q10 0.5, like nx_attention_test) 52 if nx_attn_forward(q, k, v, out_std, scale_q10) != NX_ATTN_OK { return 1 } 53 54 // ===== Flash attention with block size 2 ==================== 55 if nx_fa_forward(q, k, v, out_flash, scale_q10, 2) != NX_FA_OK { return 2 } 56 57 // ===== Oracle: epsilon-rel match (algorithmic equivalence) == 58 // 59 // The two implementations should agree within Q10 rounding 60 // tolerance. Online softmax accumulation order vs all-at-once 61 // softmax produces near-identical results; eps_q10 = 100 (~10%) 62 // is a safety margin for the small test case. 63 let witness: *i64 = (sys_mmap(NX_NO_WITNESS_FIELDS * 8)) as *i64 64 let v_eps: nx_int = nx_no_check_epsilon_rel_q10(out_flash, out_std, 100, 65 witness) 66 if v_eps != NX_NO_VERDICT_EQUAL { return 10 + v_eps } 67 68 // ===== Block size 1 (degenerate, every KV is its own block) = 69 let out_b1: *NxTensor = nx_t_alloc(NX_DT_I64, sh_out, 2, err) 70 if nx_fa_forward(q, k, v, out_b1, scale_q10, 1) != NX_FA_OK { return 20 } 71 let v_b1: nx_int = nx_no_check_epsilon_rel_q10(out_b1, out_std, 100, 72 witness) 73 if v_b1 != NX_NO_VERDICT_EQUAL { return 30 + v_b1 } 74 75 // ===== Block size 4 (everything in one block) ============== 76 let out_b4: *NxTensor = nx_t_alloc(NX_DT_I64, sh_out, 2, err) 77 if nx_fa_forward(q, k, v, out_b4, scale_q10, 4) != NX_FA_OK { return 40 } 78 let v_b4: nx_int = nx_no_check_epsilon_rel_q10(out_b4, out_std, 100, 79 witness) 80 if v_b4 != NX_NO_VERDICT_EQUAL { return 50 + v_b4 } 81 82 // ===== Memory accounting ==================================== 83 // 84 // For n_q=2, n_kv=4, d=3, block_size=2: 85 // standard: 2 * 4 * 8 = 64 bytes (full S matrix) 86 // flash: 2 * 3 * 8 + 2 * 8 * 2 = 80 bytes 87 // FlashAttn LOSES memory on tiny cases (per-row + scratch 88 // overhead > skipped S matrix). Win flips for n_kv > d * 89 // some constant. Honest test of the formula: 90 let mem_buf: *i64 = (sys_mmap(16)) as *i64 91 nx_fa_working_memory_bytes(2, 4, 3, 2, mem_buf) 92 if mem_buf[0] != 64 { return 60 } 93 if mem_buf[1] != 80 { return 61 } 94 95 // For n_q=128, n_kv=1024, d=64, block_size=64: 96 // standard: 128 * 1024 * 8 = 1 048 576 bytes (~1 MB) 97 // flash: 128 * 64 * 8 + 64 * 8 * 2 = 66560 bytes (~65 KB) 98 // ratio = 1048576 / 66560 = 15.75x 99 nx_fa_working_memory_bytes(128, 1024, 64, 64, mem_buf) 100 if mem_buf[0] != 1048576 { return 70 } 101 if mem_buf[1] != 66560 { return 71 } 102 // savings_ratio_q10 = (1048576 * 1024) / 66560 = 1073741824 / 66560 = 16131 (integer floor). 103 // (The prior 16128 was a hand-approximation from ratio ~15.75x; the exact ratio is 15.7536x.) 104 let savings: nx_int = nx_fa_savings_ratio_q10(128, 1024, 64, 64) 105 if savings != 16131 { return 72 } 106 107 // For n_q=1024, n_kv=4096, d=128, block_size=128 (realistic Z-Image scale): 108 // standard: 1024 * 4096 * 8 = 33 554 432 bytes (32 MB per head!) 109 // flash: 1024 * 128 * 8 + 128 * 8 * 2 = 1 050 624 bytes (~1 MB) 110 // ratio = ~32x 111 nx_fa_working_memory_bytes(1024, 4096, 128, 128, mem_buf) 112 if mem_buf[0] != 33554432 { return 80 } 113 if mem_buf[1] != 1050624 { return 81 } 114 115 // ===== Sealed verdict band ================================== 116 if nx_fa_verdict_is_valid(NX_FA_OK) != 1 { return 90 } 117 if nx_fa_verdict_is_valid(NX_FA_N_VERDICTS) != 0 { return 91 } 118 if nx_fa_verdict_is_valid(0 - 1) != 0 { return 92 } 119 120 return 0 121}