nx_flashattn_scale.nx source
↩ module page · 126 lines · 4977 B
1// nx_flashattn_scale.nx -- SOVEREIGN general-n FlashAttention, verified vs direct softmax on synthetic data.
2//
3// Proves the online-softmax recurrence scales to ARBITRARY n with O(d) memory: for each query we stream ALL
4// n keys through a running (max m, normalizer l, accumulator acc[d]), rescaling acc by exp(m_old-m_new).
5// The n x n score matrix is never built. Verifies bit-for-bit (within f32 tol) against a direct-softmax
6// reference on the same synthetic Q/K/V. n=64,d=32 here; identical code runs at n=4096 with the SAME O(d) mem.
7// license_tier: ORIGINAL
8import "nx_syscalls.nx"
9import "nx_le.nx"
10import "nx_f32.nx"
11import "nx_f32_div.nx"
12import "nx_f32_cvt.nx"
13import "nx_f32_exp.nx"
14import "nx_strconv.nx"
15
16func fa_seed(i: i64, dd: i64, a: i64, b: i64, m: i64) -> i64 {
17 var x: i64 = i * a + dd * b
18 x = x - (x / m) * m // x mod m, x>=0
19 return nx_f32_div(nx_i32_to_f32(x - m / 2), nx_i32_to_f32(m)) // ~[-0.5,0.5]
20}
21
22func main() -> i64 {
23 let N: i64 = 64
24 let DD: i64 = 32
25 let q: *i64 = sys_mmap(N * DD * 8) as *i64
26 let k: *i64 = sys_mmap(N * DD * 8) as *i64
27 let v: *i64 = sys_mmap(N * DD * 8) as *i64
28 var i: i64 = 0
29 while i < N {
30 var dd: i64 = 0
31 while dd < DD {
32 q[i * DD + dd] = fa_seed(i, dd, 7, 3, 13)
33 k[i * DD + dd] = fa_seed(i, dd, 5, 2, 11)
34 v[i * DD + dd] = fa_seed(i, dd, 3, 1, 9)
35 dd = dd + 1
36 }
37 i = i + 1
38 }
39
40 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(DD)))
41 let oref: *i64 = sys_mmap(N * DD * 8) as *i64
42 let ofa: *i64 = sys_mmap(N * DD * 8) as *i64
43 let sc: *i64 = sys_mmap(N * 8) as *i64
44 let acc: *i64 = sys_mmap(DD * 8) as *i64
45
46 // ---- DIRECT softmax reference (materializes the n scores) ----
47 i = 0
48 while i < N {
49 var j: i64 = 0
50 while j < N {
51 var s: i64 = 0
52 var dd: i64 = 0
53 while dd < DD { s = nx_f32_add(s, nx_f32_mul(q[i * DD + dd], k[j * DD + dd])); dd = dd + 1 }
54 sc[j] = nx_f32_mul(s, scale)
55 j = j + 1
56 }
57 var mx: i64 = sc[0]
58 j = 1
59 while j < N { if nx_f32_lt(mx, sc[j]) == 1 { mx = sc[j] } j = j + 1 }
60 var sm: i64 = 0
61 j = 0
62 while j < N { let e: i64 = nx_f32_exp(nx_f32_sub(sc[j], mx)); sc[j] = e; sm = nx_f32_add(sm, e); j = j + 1 }
63 var dd: i64 = 0
64 while dd < DD {
65 var a: i64 = 0
66 j = 0
67 while j < N { a = nx_f32_add(a, nx_f32_mul(sc[j], v[j * DD + dd])); j = j + 1 }
68 oref[i * DD + dd] = nx_f32_div(a, sm)
69 dd = dd + 1
70 }
71 i = i + 1
72 }
73
74 // ---- FLASHATTENTION streaming (O(d) memory, no n x n matrix) ----
75 i = 0
76 while i < N {
77 // key 0 seeds the running state
78 var s: i64 = 0
79 var dd: i64 = 0
80 while dd < DD { s = nx_f32_add(s, nx_f32_mul(q[i * DD + dd], k[dd])); dd = dd + 1 }
81 var m: i64 = nx_f32_mul(s, scale)
82 var l: i64 = nx_i32_to_f32(1)
83 dd = 0
84 while dd < DD { acc[dd] = v[dd]; dd = dd + 1 }
85 var j: i64 = 1
86 while j < N {
87 s = 0
88 dd = 0
89 while dd < DD { s = nx_f32_add(s, nx_f32_mul(q[i * DD + dd], k[j * DD + dd])); dd = dd + 1 }
90 s = nx_f32_mul(s, scale)
91 var mnew: i64 = m
92 if nx_f32_lt(mnew, s) == 1 { mnew = s }
93 let corr: i64 = nx_f32_exp(nx_f32_sub(m, mnew))
94 let p: i64 = nx_f32_exp(nx_f32_sub(s, mnew))
95 l = nx_f32_add(nx_f32_mul(l, corr), p)
96 dd = 0
97 while dd < DD { acc[dd] = nx_f32_add(nx_f32_mul(acc[dd], corr), nx_f32_mul(p, v[j * DD + dd])); dd = dd + 1 }
98 m = mnew
99 j = j + 1
100 }
101 dd = 0
102 while dd < DD { ofa[i * DD + dd] = nx_f32_div(acc[dd], l); dd = dd + 1 }
103 i = i + 1
104 }
105
106 // ---- compare ----
107 let tol: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000))
108 var fails: i64 = 0
109 i = 0
110 while i < N * DD {
111 if (nx_f32_sub(oref[i], ofa[i]) & 0x7FFFFFFF) >= tol { fails = fails + 1 }
112 i = i + 1
113 }
114 let ofd: i64 = sys_openat_wr("/tmp/fasc.txt" as *u8, 0x1a4)
115 if ofd >= 0 {
116 let dec: *u8 = sys_mmap(32)
117 sys_write(ofd, "n=" as *u8, 2); let na: i64 = nx_strconv_format_i64(N, dec); sys_write(ofd, dec, na)
118 sys_write(ofd, " d=" as *u8, 3); let nb: i64 = nx_strconv_format_i64(DD, dec); sys_write(ofd, dec, nb)
119 sys_write(ofd, " fails=" as *u8, 7); let nc: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, nc)
120 sys_write(ofd, " (FA_mem=O(d)=" as *u8, 14); let nd2: i64 = nx_strconv_format_i64(DD, dec); sys_write(ofd, dec, nd2)
121 sys_write(ofd, " vs direct=O(n*n)=" as *u8, 18); let ne: i64 = nx_strconv_format_i64(N * N, dec); sys_write(ofd, dec, ne)
122 sys_write(ofd, ")\n" as *u8, 2); sys_close(ofd)
123 }
124 if fails > 0 { return 20 }
125 return 0
126}