code wiki / (root) / nx_zimage_flashattn_verify.nx

nx_zimage_flashattn_verify.nx source

↩ module page · 155 lines · 6751 B

1// nx_zimage_flashattn_verify.nx -- SOVEREIGN FlashAttention (online-softmax) SDPA, verified EXACT vs oracle. 2// 3// The O(n)-memory attention: instead of building the n x n score matrix, softmaxing, then weighting V, 4// we STREAM keys into a running accumulator, rescaling it by exp(m_old - m_new) as the running max grows. 5// Peak memory is O(d) per query (the accumulator + 3 scalars m,l,a), NOT O(n) -- so it scales to any n 6// without materializing scores. This organ proves the recurrence is NUMERICALLY IDENTICAL to standard 7// softmax attention (verifies vs the same oracle ao_pre as the direct SDPA). Lossless, not an approximation. 8// Uses a LOCAL accumulator (no mmap scratch in the loop) per the nx_cc codegen lesson. No 3rd party. 9// license_tier: ORIGINAL 10import "nx_syscalls.nx" 11import "nx_le.nx" 12import "nx_f32.nx" 13import "nx_f32_div.nx" 14import "nx_f32_cvt.nx" 15import "nx_f32_exp.nx" 16import "nx_strconv.nx" 17const K_MAGIC_3840: i64 = 3840 18const K_MAGIC_11520: i64 = 11520 19 20func zfa_load(name: *u8, nl: i64, n_floats: i64) -> *u8 { 21 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8 22 let path: *u8 = sys_mmap(256) 23 var p: i64 = 0 24 var i: i64 = 0 25 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 } 26 i = 0 27 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 } 28 path[p] = 0x2E; p = p + 1 29 path[p] = 0x66; p = p + 1 30 path[p] = 0x33; p = p + 1 31 path[p] = 0x32; p = p + 1 32 path[p] = 0 33 let fd: i64 = sys_openat_rd(path) 34 if fd < 0 { return 0 as *u8 } 35 let bytes: i64 = n_floats * 4 36 let buf: *u8 = sys_mmap(bytes + 64) 37 var tot: i64 = 0 38 var go: i64 = 1 39 while go == 1 { let r: i64 = sys_read(fd, ((buf as i64) + tot) as *u8, bytes - tot); if r <= 0 { go = 0 } else { tot = tot + r; if tot >= bytes { go = 0 } } } 40 sys_close(fd) 41 return buf 42} 43 44func main() -> i64 { 45 let D: i64 = K_MAGIC_3840 46 let H: i64 = 30 47 let HD: i64 = 128 48 let NT: i64 = 4 49 let QKV: i64 = K_MAGIC_11520 50 let qr: *u8 = zfa_load("q_roped" as *u8, 7, NT * D) 51 let kr: *u8 = zfa_load("k_roped" as *u8, 7, NT * D) 52 let qkvd: *u8 = zfa_load("qkv" as *u8, 3, NT * QKV) 53 let aog: *u8 = zfa_load("ao_pre" as *u8, 6, NT * D) 54 if (qr as i64) == 0 { return 30 } 55 if (kr as i64) == 0 { return 32 } 56 if (qkvd as i64) == 0 { return 33 } 57 if (aog as i64) == 0 { return 31 } 58 59 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(128))) 60 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100)) 61 let one: i64 = nx_i32_to_f32(1) 62 let vbase: i64 = D + D 63 var fails: i64 = 0 64 var first_bad: i64 = 0 - 1 65 66 var hd: i64 = 0 67 while hd < H { 68 var t1: i64 = 0 69 while t1 < NT { 70 // --- scores q.k for the 4 keys (one d-loop) --- 71 let qb: i64 = (t1 * D + hd * HD) * 4 72 let k0b: i64 = (hd * HD) * 4 73 let k1b: i64 = (D + hd * HD) * 4 74 let k2b: i64 = (2 * D + hd * HD) * 4 75 let k3b: i64 = (3 * D + hd * HD) * 4 76 var s0: i64 = 0 77 var s1: i64 = 0 78 var s2: i64 = 0 79 var s3: i64 = 0 80 var d: i64 = 0 81 while d < HD { 82 let qd: i64 = nx_le_read_u32(qr, qb + d * 4) 83 s0 = nx_f32_add(s0, nx_f32_mul(qd, nx_le_read_u32(kr, k0b + d * 4))) 84 s1 = nx_f32_add(s1, nx_f32_mul(qd, nx_le_read_u32(kr, k1b + d * 4))) 85 s2 = nx_f32_add(s2, nx_f32_mul(qd, nx_le_read_u32(kr, k2b + d * 4))) 86 s3 = nx_f32_add(s3, nx_f32_mul(qd, nx_le_read_u32(kr, k3b + d * 4))) 87 d = d + 1 88 } 89 s0 = nx_f32_mul(s0, scale) 90 s1 = nx_f32_mul(s1, scale) 91 s2 = nx_f32_mul(s2, scale) 92 s3 = nx_f32_mul(s3, scale) 93 94 // --- ONLINE SOFTMAX: streaming rescale scalars (never store all scores) --- 95 // key0: m=s0, l=1, acc=v0 96 var m: i64 = s0 97 // key1 98 var mn: i64 = m 99 if nx_f32_lt(mn, s1) == 1 { mn = s1 } 100 let corr0: i64 = nx_f32_exp(nx_f32_sub(m, mn)) 101 let p1: i64 = nx_f32_exp(nx_f32_sub(s1, mn)) 102 m = mn 103 // key2 104 mn = m 105 if nx_f32_lt(mn, s2) == 1 { mn = s2 } 106 let corr1: i64 = nx_f32_exp(nx_f32_sub(m, mn)) 107 let p2: i64 = nx_f32_exp(nx_f32_sub(s2, mn)) 108 m = mn 109 // key3 110 mn = m 111 if nx_f32_lt(mn, s3) == 1 { mn = s3 } 112 let corr2: i64 = nx_f32_exp(nx_f32_sub(m, mn)) 113 let p3: i64 = nx_f32_exp(nx_f32_sub(s3, mn)) 114 m = mn 115 // running normalizer l (same recurrence) 116 var l: i64 = one 117 l = nx_f32_add(nx_f32_mul(l, corr0), p1) 118 l = nx_f32_add(nx_f32_mul(l, corr1), p2) 119 l = nx_f32_add(nx_f32_mul(l, corr2), p3) 120 121 let v0b: i64 = (vbase + hd * HD) * 4 122 let v1b: i64 = (QKV + vbase + hd * HD) * 4 123 let v2b: i64 = (2 * QKV + vbase + hd * HD) * 4 124 let v3b: i64 = (3 * QKV + vbase + hd * HD) * 4 125 126 // --- per-channel online accumulate with a LOCAL (no scratch array) --- 127 var d2: i64 = 0 128 while d2 < HD { 129 var a: i64 = nx_le_read_u32(qkvd, v0b + d2 * 4) // acc_0 = v0 130 a = nx_f32_add(nx_f32_mul(a, corr0), nx_f32_mul(p1, nx_le_read_u32(qkvd, v1b + d2 * 4))) // key1 131 a = nx_f32_add(nx_f32_mul(a, corr1), nx_f32_mul(p2, nx_le_read_u32(qkvd, v2b + d2 * 4))) // key2 132 a = nx_f32_add(nx_f32_mul(a, corr2), nx_f32_mul(p3, nx_le_read_u32(qkvd, v3b + d2 * 4))) // key3 133 a = nx_f32_div(a, l) 134 let g: i64 = nx_le_read_u32(aog, (t1 * D + hd * HD + d2) * 4) 135 var thr: i64 = tolc 136 let ag: i64 = g & 0x7FFFFFFF 137 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) } 138 if (nx_f32_sub(a, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t1 * D + hd * HD + d2 } } 139 d2 = d2 + 1 140 } 141 t1 = t1 + 1 142 } 143 hd = hd + 1 144 } 145 146 let ofd: i64 = sys_openat_wr("/tmp/zfa.txt" as *u8, 0x1a4) 147 if ofd >= 0 { 148 let dec: *u8 = sys_mmap(32) 149 sys_write(ofd, "fails=" as *u8, 6); let n1: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, n1) 150 sys_write(ofd, " first_bad=" as *u8, 11); let n2: i64 = nx_strconv_format_i64(first_bad, dec); sys_write(ofd, dec, n2) 151 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd) 152 } 153 if fails > 0 { return 20 } 154 return 0 155}