code wiki / (root) / nx_zimage_sdpaonly_verify.nx

nx_zimage_sdpaonly_verify.nx source

↩ module page · 131 lines · 5644 B

1// nx_zimage_sdpaonly_verify.nx -- SOVEREIGN Z-Image attention sub-piece 2c: SDPA, verified vs oracle. 2// Rewritten with NT=4 UNROLLED locals (no sc[] scratch, no nested t2 loop) to dodge the codegen crash. 3// Reads oracle q_roped/k_roped + qkv(V) -> per-head softmax(q.k/sqrt128) v -> ao_pre. 4// license_tier: ORIGINAL 5import "nx_syscalls.nx" 6import "nx_le.nx" 7import "nx_f32.nx" 8import "nx_f32_div.nx" 9import "nx_f32_cvt.nx" 10import "nx_f32_exp.nx" 11import "nx_strconv.nx" 12const K_MAGIC_3840: i64 = 3840 13const K_MAGIC_11520: i64 = 11520 14const K_MAGIC_7680: i64 = 7680 15 16func zso_load(name: *u8, nl: i64, n_floats: i64) -> *u8 { 17 let base: *u8 = "/mnt/c/Users/elder/AppData/Local/Temp/claude/C--Users-elder/7be78b15-304c-449e-afe8-4d5bd7ddaa9c/scratchpad/zblk/" as *u8 18 let path: *u8 = sys_mmap(256) 19 var p: i64 = 0 20 var i: i64 = 0 21 while base[i] != 0 { path[p] = base[i]; p = p + 1; i = i + 1 } 22 i = 0 23 while i < nl { path[p] = name[i]; p = p + 1; i = i + 1 } 24 path[p] = 0x2E; p = p + 1 25 path[p] = 0x66; p = p + 1 26 path[p] = 0x33; p = p + 1 27 path[p] = 0x32; p = p + 1 28 path[p] = 0 29 let fd: i64 = sys_openat_rd(path) 30 if fd < 0 { return 0 as *u8 } 31 let bytes: i64 = n_floats * 4 32 let buf: *u8 = sys_mmap(bytes + 64) 33 var tot: i64 = 0 34 var go: i64 = 1 35 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 } } } 36 sys_close(fd) 37 return buf 38} 39 40func main() -> i64 { 41 let D: i64 = K_MAGIC_3840 42 let H: i64 = 30 43 let HD: i64 = 128 44 let NT: i64 = 4 45 let QKV: i64 = K_MAGIC_11520 46 let qr: *u8 = zso_load("q_roped" as *u8, 7, NT * D) 47 let kr: *u8 = zso_load("k_roped" as *u8, 7, NT * D) 48 let qkvd: *u8 = zso_load("qkv" as *u8, 3, NT * QKV) 49 let aog: *u8 = zso_load("ao_pre" as *u8, 6, NT * D) 50 if (qr as i64) == 0 { return 30 } 51 if (kr as i64) == 0 { return 32 } 52 if (qkvd as i64) == 0 { return 33 } 53 if (aog as i64) == 0 { return 31 } 54 55 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(128))) 56 let tolc: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100)) 57 let vbase: i64 = D + D // V region offset = K_MAGIC_7680 58 var fails: i64 = 0 59 var first_bad: i64 = 0 - 1 60 61 var hd: i64 = 0 62 while hd < H { 63 var t1: i64 = 0 64 while t1 < NT { 65 let qb: i64 = (t1 * D + hd * HD) * 4 66 let k0b: i64 = (hd * HD) * 4 67 let k1b: i64 = (D + hd * HD) * 4 68 let k2b: i64 = (2 * D + hd * HD) * 4 69 let k3b: i64 = (3 * D + hd * HD) * 4 70 var s0: i64 = 0 71 var s1: i64 = 0 72 var s2: i64 = 0 73 var s3: i64 = 0 74 var d: i64 = 0 75 while d < HD { 76 let qd: i64 = nx_le_read_u32(qr, qb + d * 4) 77 s0 = nx_f32_add(s0, nx_f32_mul(qd, nx_le_read_u32(kr, k0b + d * 4))) 78 s1 = nx_f32_add(s1, nx_f32_mul(qd, nx_le_read_u32(kr, k1b + d * 4))) 79 s2 = nx_f32_add(s2, nx_f32_mul(qd, nx_le_read_u32(kr, k2b + d * 4))) 80 s3 = nx_f32_add(s3, nx_f32_mul(qd, nx_le_read_u32(kr, k3b + d * 4))) 81 d = d + 1 82 } 83 s0 = nx_f32_mul(s0, scale) 84 s1 = nx_f32_mul(s1, scale) 85 s2 = nx_f32_mul(s2, scale) 86 s3 = nx_f32_mul(s3, scale) 87 var mx: i64 = s0 88 if nx_f32_lt(mx, s1) == 1 { mx = s1 } 89 if nx_f32_lt(mx, s2) == 1 { mx = s2 } 90 if nx_f32_lt(mx, s3) == 1 { mx = s3 } 91 let e0: i64 = nx_f32_exp(nx_f32_sub(s0, mx)) 92 let e1: i64 = nx_f32_exp(nx_f32_sub(s1, mx)) 93 let e2: i64 = nx_f32_exp(nx_f32_sub(s2, mx)) 94 let e3: i64 = nx_f32_exp(nx_f32_sub(s3, mx)) 95 let sum: i64 = nx_f32_add(nx_f32_add(e0, e1), nx_f32_add(e2, e3)) 96 let a0: i64 = nx_f32_div(e0, sum) 97 let a1: i64 = nx_f32_div(e1, sum) 98 let a2: i64 = nx_f32_div(e2, sum) 99 let a3: i64 = nx_f32_div(e3, sum) 100 let v0b: i64 = (vbase + hd * HD) * 4 101 let v1b: i64 = (QKV + vbase + hd * HD) * 4 102 let v2b: i64 = (2 * QKV + vbase + hd * HD) * 4 103 let v3b: i64 = (3 * QKV + vbase + hd * HD) * 4 104 var d2: i64 = 0 105 while d2 < HD { 106 var a: i64 = nx_f32_mul(a0, nx_le_read_u32(qkvd, v0b + d2 * 4)) 107 a = nx_f32_add(a, nx_f32_mul(a1, nx_le_read_u32(qkvd, v1b + d2 * 4))) 108 a = nx_f32_add(a, nx_f32_mul(a2, nx_le_read_u32(qkvd, v2b + d2 * 4))) 109 a = nx_f32_add(a, nx_f32_mul(a3, nx_le_read_u32(qkvd, v3b + d2 * 4))) 110 let g: i64 = nx_le_read_u32(aog, (t1 * D + hd * HD + d2) * 4) 111 var thr: i64 = tolc 112 let ag: i64 = g & 0x7FFFFFFF 113 if nx_f32_lt(thr, nx_f32_mul(tolc, ag)) == 1 { thr = nx_f32_mul(tolc, ag) } 114 if (nx_f32_sub(a, g) & 0x7FFFFFFF) >= thr { fails = fails + 1; if first_bad < 0 { first_bad = t1 * D + hd * HD + d2 } } 115 d2 = d2 + 1 116 } 117 t1 = t1 + 1 118 } 119 hd = hd + 1 120 } 121 122 let ofd: i64 = sys_openat_wr("/tmp/zso.txt" as *u8, 0x1a4) 123 if ofd >= 0 { 124 let dec: *u8 = sys_mmap(32) 125 sys_write(ofd, "fails=" as *u8, 6); let n1: i64 = nx_strconv_format_i64(fails, dec); sys_write(ofd, dec, n1) 126 sys_write(ofd, " first_bad=" as *u8, 11); let n2: i64 = nx_strconv_format_i64(first_bad, dec); sys_write(ofd, dec, n2) 127 sys_write(ofd, "\n" as *u8, 1); sys_close(ofd) 128 } 129 if fails > 0 { return 20 } 130 return 0 131}