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}