code wiki / (root) / nx_f32_dit_block_linear.nx

nx_f32_dit_block_linear.nx source

↩ module page · 191 lines · 6969 B

1// nx_f32_dit_block_linear.nx -- SANA-style DiT block using LINEAR attention (the efficient Z-Image DiT unit). 2// 3// From DiT-IC/SANA (ingested 2026-07-01): a DiT block with O(n) linear attention instead of O(n^2) softmax. 4// adaLN(x; sc1,sh1) -> Q/K/V -> linear-attn -> Wo -> x1 = x + gt1*ao 5// adaLN(x1; sc2,sh2) -> SwiGLU(Wg,Wu,Wd) -> out = x1 + gt2*ffn 6// adaLN-Zero gate (gt1=gt2=0) => out == x bit-exact, proving the WIRING independent of attn/FFN internals 7// (linear-attn + SwiGLU are separately gated). Single-head (head_dim=D). Linear-attn inlined; SiLU via lib. 8// license_tier: ORIGINAL 9import "nx_syscalls.nx" 10import "nx_f32.nx" 11import "nx_f32_div.nx" 12import "nx_f32_cvt.nx" 13import "nx_f32_activations.nx" 14 15func dbl_phi(v: i64) -> i64 { 16 var r: i64 = v 17 if (v & 0x80000000) != 0 { r = 0 } 18 return nx_f32_add(r, nx_i32_to_f32(1)) 19} 20 21// out[t*od+o] = sum_i in[t*id+i] * W[o*id+i] 22func dbl_linear(inp: *i64, W: *i64, out: *i64, n: i64, id: i64, od: i64) -> i64 { 23 var t: i64 = 0 24 while t < n { 25 var o: i64 = 0 26 while o < od { 27 var acc: i64 = 0 28 var i: i64 = 0 29 while i < id { acc = nx_f32_add(acc, nx_f32_mul(inp[t * id + i], W[o * id + i])); i = i + 1 } 30 out[t * od + o] = acc 31 o = o + 1 32 } 33 t = t + 1 34 } 35 return 0 36} 37 38// adaLN modulate: out[t*D+d] = x[t*D+d]*(1+sc[d]) + sh[d] 39func dbl_modulate(x: *i64, sc: *i64, sh: *i64, out: *i64, n: i64, D: i64) -> i64 { 40 var t: i64 = 0 41 while t < n { 42 var d: i64 = 0 43 while d < D { 44 out[t * D + d] = nx_f32_add(nx_f32_mul(x[t * D + d], nx_f32_add(nx_i32_to_f32(1), sc[d])), sh[d]) 45 d = d + 1 46 } 47 t = t + 1 48 } 49 return 0 50} 51 52func dbl_linattn(Q: *i64, K: *i64, V: *i64, n: i64, D: i64, out: *i64, S: *i64, z: *i64) -> i64 { 53 var a: i64 = 0 54 while a < D * D { S[a] = 0; a = a + 1 } 55 a = 0 56 while a < D { z[a] = 0; a = a + 1 } 57 var j: i64 = 0 58 while j < n { 59 a = 0 60 while a < D { 61 let pk: i64 = dbl_phi(K[j * D + a]) 62 z[a] = nx_f32_add(z[a], pk) 63 var b: i64 = 0 64 while b < D { S[a * D + b] = nx_f32_add(S[a * D + b], nx_f32_mul(pk, V[j * D + b])); b = b + 1 } 65 a = a + 1 66 } 67 j = j + 1 68 } 69 var i: i64 = 0 70 while i < n { 71 var den: i64 = 0 72 a = 0 73 while a < D { den = nx_f32_add(den, nx_f32_mul(dbl_phi(Q[i * D + a]), z[a])); a = a + 1 } 74 var b: i64 = 0 75 while b < D { 76 var num: i64 = 0 77 a = 0 78 while a < D { num = nx_f32_add(num, nx_f32_mul(dbl_phi(Q[i * D + a]), S[a * D + b])); a = a + 1 } 79 out[i * D + b] = nx_f32_div(num, den) 80 b = b + 1 81 } 82 i = i + 1 83 } 84 return 0 85} 86 87// gated residual: x1[k] = x[k] + gt[k mod D] * y[k] 88func dbl_gate_res(x: *i64, y: *i64, gt: *i64, out: *i64, n: i64, D: i64) -> i64 { 89 var t: i64 = 0 90 while t < n { 91 var d: i64 = 0 92 while d < D { out[t * D + d] = nx_f32_add(x[t * D + d], nx_f32_mul(gt[d], y[t * D + d])); d = d + 1 } 93 t = t + 1 94 } 95 return 0 96} 97 98func nx_f32_dit_block_linear(x: *i64, n: i64, D: i64, d_ff: i64, 99 Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, sc1: *i64, sh1: *i64, gt1: *i64, 100 Wg: *i64, Wu: *i64, Wd: *i64, sc2: *i64, sh2: *i64, gt2: *i64, out: *i64) -> i64 { 101 let hn: *i64 = sys_mmap(n * D * 8) as *i64 102 let Q: *i64 = sys_mmap(n * D * 8) as *i64 103 let K: *i64 = sys_mmap(n * D * 8) as *i64 104 let V: *i64 = sys_mmap(n * D * 8) as *i64 105 let attn: *i64 = sys_mmap(n * D * 8) as *i64 106 let ao: *i64 = sys_mmap(n * D * 8) as *i64 107 let x1: *i64 = sys_mmap(n * D * 8) as *i64 108 let S: *i64 = sys_mmap(D * D * 8) as *i64 109 let z: *i64 = sys_mmap(D * 8) as *i64 110 111 // attention sublayer 112 dbl_modulate(x, sc1, sh1, hn, n, D) 113 dbl_linear(hn, Wq, Q, n, D, D) 114 dbl_linear(hn, Wk, K, n, D, D) 115 dbl_linear(hn, Wv, V, n, D, D) 116 dbl_linattn(Q, K, V, n, D, attn, S, z) 117 dbl_linear(attn, Wo, ao, n, D, D) 118 dbl_gate_res(x, ao, gt1, x1, n, D) 119 120 // FFN sublayer (SwiGLU) 121 let hn2: *i64 = sys_mmap(n * D * 8) as *i64 122 let hg: *i64 = sys_mmap(n * d_ff * 8) as *i64 123 let hu: *i64 = sys_mmap(n * d_ff * 8) as *i64 124 let hm: *i64 = sys_mmap(n * d_ff * 8) as *i64 125 let ffn: *i64 = sys_mmap(n * D * 8) as *i64 126 dbl_modulate(x1, sc2, sh2, hn2, n, D) 127 dbl_linear(hn2, Wg, hg, n, D, d_ff) 128 dbl_linear(hn2, Wu, hu, n, D, d_ff) 129 var t: i64 = 0 130 while t < n { 131 var f: i64 = 0 132 while f < d_ff { hm[t * d_ff + f] = nx_f32_mul(nx_f32_silu(hg[t * d_ff + f]), hu[t * d_ff + f]); f = f + 1 } 133 t = t + 1 134 } 135 dbl_linear(hm, Wd, ffn, n, d_ff, D) 136 dbl_gate_res(x1, ffn, gt2, out, n, D) 137 return 0 138} 139 140func main() -> i64 { 141 let n: i64 = 2 142 let D: i64 = 4 143 let dff: i64 = 4 144 let x: *i64 = sys_mmap(n * D * 8) as *i64 145 let out: *i64 = sys_mmap(n * D * 8) as *i64 146 let Wq: *i64 = sys_mmap(D * D * 8) as *i64 147 let Wk: *i64 = sys_mmap(D * D * 8) as *i64 148 let Wv: *i64 = sys_mmap(D * D * 8) as *i64 149 let Wo: *i64 = sys_mmap(D * D * 8) as *i64 150 let Wg: *i64 = sys_mmap(dff * D * 8) as *i64 151 let Wu: *i64 = sys_mmap(dff * D * 8) as *i64 152 let Wd: *i64 = sys_mmap(D * dff * 8) as *i64 153 let sc1: *i64 = sys_mmap(D * 8) as *i64 154 let sh1: *i64 = sys_mmap(D * 8) as *i64 155 let gt1: *i64 = sys_mmap(D * 8) as *i64 156 let sc2: *i64 = sys_mmap(D * 8) as *i64 157 let sh2: *i64 = sys_mmap(D * 8) as *i64 158 let gt2: *i64 = sys_mmap(D * 8) as *i64 159 let p1: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(10)) 160 161 var i: i64 = 0 162 while i < n * D { x[i] = nx_i32_to_f32(i + 1); i = i + 1 } 163 i = 0 164 while i < D * D { Wq[i] = p1; Wk[i] = p1; Wv[i] = p1; Wo[i] = p1; i = i + 1 } 165 i = 0 166 while i < dff * D { Wg[i] = p1; Wu[i] = p1; Wd[i] = p1; i = i + 1 } 167 // adaLN scale/shift arbitrary; gates ZERO -> identity 168 i = 0 169 while i < D { sc1[i] = p1; sh1[i] = p1; sc2[i] = p1; sh2[i] = p1; gt1[i] = 0; gt2[i] = 0; i = i + 1 } 170 171 nx_f32_dit_block_linear(x, n, D, dff, Wq, Wk, Wv, Wo, sc1, sh1, gt1, Wg, Wu, Wd, sc2, sh2, gt2, out) 172 // adaLN-Zero: gt=0 -> out == x bit-exact 173 i = 0 174 while i < n * D { if out[i] != x[i] { return 10 } i = i + 1 } 175 176 // non-zero gates -> block transforms (out != x somewhere) + deterministic 177 i = 0 178 while i < D { gt1[i] = p1; gt2[i] = p1; i = i + 1 } 179 let out2: *i64 = sys_mmap(n * D * 8) as *i64 180 nx_f32_dit_block_linear(x, n, D, dff, Wq, Wk, Wv, Wo, sc1, sh1, gt1, Wg, Wu, Wd, sc2, sh2, gt2, out) 181 nx_f32_dit_block_linear(x, n, D, dff, Wq, Wk, Wv, Wo, sc1, sh1, gt1, Wg, Wu, Wd, sc2, sh2, gt2, out2) 182 var diff: i64 = 0 183 i = 0 184 while i < n * D { 185 if out[i] != out2[i] { return 20 } // deterministic 186 if out[i] != x[i] { diff = 1 } 187 i = i + 1 188 } 189 if diff == 0 { return 21 } 190 return 0 191}