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}