code wiki / (root) / nx_flashattn_scale.nx

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}