nx_f32_dit_block_tiny.nx source
↩ module page · 169 lines · 7403 B
1// nx_f32_dit_block_tiny.nx -- first sovereign f32 DiT (Diffusion Transformer) BLOCK, assembled from the
2// gated bricks. The repeating unit of the Z-Image diffusion backbone (Peebles&Xie 2023 DiT; adaLN-Zero).
3//
4// sd-server -> Nishi migration milestone (the DiT half's analogue of nx_f32_vae_decode_tiny). Per token:
5// x1 = x + gate1 * Wo @ Attention( Wq,Wk,Wv @ adaLN_mod( RMSNorm(x), scale1, shift1 ) )
6// out = x1 + gate2 * SwiGLU( adaLN_mod( RMSNorm(x1), scale2, shift2 ) )
7// Composes the gated organs: nx_f32_rmsnorm + nx_f32_adaln (modulate + gated residual) + nx_f32_attention +
8// nx_f32_silu (in the inline SwiGLU) + inline linear projections. Single head v1 (head_dim = D). The adaLN
9// scale/shift/gate come from the conditioning MLP (caller-supplied; adaLN-Zero starts gate=0 => identity).
10//
11// x,out: flat *i64 f32 bits [n_tokens, D]. W{q,k,v,o}: [D,D]. SwiGLU W1,W3: [d_ff,D], W2: [D,d_ff].
12// sc/sh/gt{1,2}: [D]. gamma: [D] (RMSNorm scale). Scratch is allocated internally (fork-per-use organ).
13// license_tier: ORIGINAL
14import "nx_syscalls.nx"
15import "nx_f32.nx"
16import "nx_f32_cvt.nx"
17import "nx_f32_rmsnorm.nx"
18import "nx_f32_attention.nx"
19import "nx_f32_adaln.nx"
20import "nx_f32_activations.nx"
21const K_MAGIC_100000: i64 = 100000
22
23// out[t][o] = sum_i in[t][i] * W[o][i] (W is [Dout, Din] row-major)
24func ditb_linear(inp: *i64, n_tokens: i64, Din: i64, W: *i64, Dout: i64, out: *i64) -> i64 {
25 var t: i64 = 0
26 while t < n_tokens {
27 var o: i64 = 0
28 while o < Dout {
29 var acc: i64 = 0
30 var i: i64 = 0
31 while i < Din { acc = nx_f32_add(acc, nx_f32_mul(inp[t * Din + i], W[o * Din + i])); i = i + 1 }
32 out[t * Dout + o] = acc
33 o = o + 1
34 }
35 t = t + 1
36 }
37 return 0
38}
39
40// SwiGLU FFN per token: g[f] = silu(in.W1[f]) * (in.W3[f]); out[d] = g.W2[d]
41func ditb_swiglu(inp: *i64, n_tokens: i64, D: i64, d_ff: i64, W1: *i64, W3: *i64, W2: *i64, out: *i64, h: *i64) -> i64 {
42 var t: i64 = 0
43 while t < n_tokens {
44 var f: i64 = 0
45 while f < d_ff {
46 var a1: i64 = 0
47 var a3: i64 = 0
48 var i: i64 = 0
49 while i < D { a1 = nx_f32_add(a1, nx_f32_mul(inp[t * D + i], W1[f * D + i])); a3 = nx_f32_add(a3, nx_f32_mul(inp[t * D + i], W3[f * D + i])); i = i + 1 }
50 h[f] = nx_f32_mul(nx_f32_silu(a1), a3)
51 f = f + 1
52 }
53 var d: i64 = 0
54 while d < D {
55 var o: i64 = 0
56 var ff: i64 = 0
57 while ff < d_ff { o = nx_f32_add(o, nx_f32_mul(h[ff], W2[d * d_ff + ff])); ff = ff + 1 }
58 out[t * D + d] = o
59 d = d + 1
60 }
61 t = t + 1
62 }
63 return 0
64}
65
66func nx_f32_dit_block_tiny(x: *i64, n_tokens: i64, D: i64, d_ff: i64, gamma: *i64, eps: i64,
67 Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, sc1: *i64, sh1: *i64, gt1: *i64,
68 W1: *i64, W3: *i64, W2: *i64, sc2: *i64, sh2: *i64, gt2: *i64,
69 out: *i64) -> i64 {
70 let nb: *i64 = sys_mmap(n_tokens * D * 8) as *i64
71 let mb: *i64 = sys_mmap(n_tokens * D * 8) as *i64
72 let Q: *i64 = sys_mmap(n_tokens * D * 8) as *i64
73 let K: *i64 = sys_mmap(n_tokens * D * 8) as *i64
74 let V: *i64 = sys_mmap(n_tokens * D * 8) as *i64
75 let ab: *i64 = sys_mmap(n_tokens * D * 8) as *i64
76 let ob: *i64 = sys_mmap(n_tokens * D * 8) as *i64
77 let fb: *i64 = sys_mmap(n_tokens * D * 8) as *i64
78 let sr: *i64 = sys_mmap(n_tokens * 8) as *i64
79 let pr: *i64 = sys_mmap(n_tokens * 8) as *i64
80 let hh: *i64 = sys_mmap(d_ff * 8) as *i64
81 let x1: *i64 = sys_mmap(n_tokens * D * 8) as *i64
82 let scale: i64 = nx_i32_to_f32(1)
83
84 // ---- attention sublayer ----
85 var t: i64 = 0
86 while t < n_tokens {
87 nx_f32_rmsnorm(((x as i64) + t * D * 8) as *i64, gamma, D, eps, ((nb as i64) + t * D * 8) as *i64)
88 t = t + 1
89 }
90 nx_f32_adaln_modulate(nb, n_tokens, D, sc1, sh1, mb)
91 ditb_linear(mb, n_tokens, D, Wq, D, Q)
92 ditb_linear(mb, n_tokens, D, Wk, D, K)
93 ditb_linear(mb, n_tokens, D, Wv, D, V)
94 nx_f32_attention(Q, K, V, n_tokens, D, scale, ab, sr, pr)
95 ditb_linear(ab, n_tokens, D, Wo, D, ob)
96 nx_f32_adaln_gate(x, ob, n_tokens, D, gt1, x1) // x1 = x + gate1 * ob
97
98 // ---- FFN sublayer ----
99 t = 0
100 while t < n_tokens {
101 nx_f32_rmsnorm(((x1 as i64) + t * D * 8) as *i64, gamma, D, eps, ((nb as i64) + t * D * 8) as *i64)
102 t = t + 1
103 }
104 nx_f32_adaln_modulate(nb, n_tokens, D, sc2, sh2, mb)
105 ditb_swiglu(mb, n_tokens, D, d_ff, W1, W3, W2, fb, hh)
106 nx_f32_adaln_gate(x1, fb, n_tokens, D, gt2, out) // out = x1 + gate2 * fb
107 return 0
108}
109
110// ===== Self-test (inline integration gate) ========================
111// (a) adaLN-ZERO INIT: gate1 == gate2 == 0 -> the block is IDENTITY -> out == x (bit-exact),
112// regardless of all other weights (proves the whole assembly composes + the gated-residual structure).
113// (b) gates == 1 with non-trivial weights -> the block CONTRIBUTES -> out != x somewhere.
114func main() -> i64 {
115 let n_tokens: i64 = 2
116 let D: i64 = 2
117 let d_ff: i64 = 2
118 let x: *i64 = sys_mmap(n_tokens * D * 8) as *i64
119 let out: *i64 = sys_mmap(n_tokens * D * 8) as *i64
120 let gamma: *i64 = sys_mmap(D * 8) as *i64
121 let Wq: *i64 = sys_mmap(D * D * 8) as *i64
122 let Wk: *i64 = sys_mmap(D * D * 8) as *i64
123 let Wv: *i64 = sys_mmap(D * D * 8) as *i64
124 let Wo: *i64 = sys_mmap(D * D * 8) as *i64
125 let W1: *i64 = sys_mmap(d_ff * D * 8) as *i64
126 let W3: *i64 = sys_mmap(d_ff * D * 8) as *i64
127 let W2: *i64 = sys_mmap(D * d_ff * 8) as *i64
128 let sc1: *i64 = sys_mmap(D * 8) as *i64
129 let sh1: *i64 = sys_mmap(D * 8) as *i64
130 let gt1: *i64 = sys_mmap(D * 8) as *i64
131 let sc2: *i64 = sys_mmap(D * 8) as *i64
132 let sh2: *i64 = sys_mmap(D * 8) as *i64
133 let gt2: *i64 = sys_mmap(D * 8) as *i64
134 let one: i64 = nx_i32_to_f32(1)
135 let p1: i64 = nx_f32_div(one, nx_i32_to_f32(10)) // 0.1 (needs nx_f32_div via rmsnorm import chain)
136 let eps: i64 = nx_f32_div(one, nx_i32_to_f32(K_MAGIC_100000))
137
138 // non-trivial weights everywhere; modulation scale/shift = 0
139 var i: i64 = 0
140 while i < D * D { Wq[i] = p1; Wk[i] = p1; Wv[i] = p1; Wo[i] = p1; i = i + 1 }
141 i = 0
142 while i < d_ff * D { W1[i] = p1; W3[i] = p1; i = i + 1 }
143 i = 0
144 while i < D * d_ff { W2[i] = p1; i = i + 1 }
145 i = 0
146 while i < D { gamma[i] = one; sc1[i] = 0; sh1[i] = 0; sc2[i] = 0; sh2[i] = 0; i = i + 1 }
147 i = 0
148 while i < n_tokens * D { x[i] = nx_i32_to_f32(i + 1); i = i + 1 }
149
150 // (a) adaLN-Zero: gates 0 -> identity
151 i = 0
152 while i < D { gt1[i] = 0; gt2[i] = 0; i = i + 1 }
153 let va: i64 = nx_f32_dit_block_tiny(x, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, sc1, sh1, gt1, W1, W3, W2, sc2, sh2, gt2, out)
154 if va != 0 { return 10 }
155 i = 0
156 while i < n_tokens * D { if out[i] != x[i] { return 20 } i = i + 1 }
157
158 // (b) gates 1 -> contributes
159 i = 0
160 while i < D { gt1[i] = one; gt2[i] = one; i = i + 1 }
161 let vb: i64 = nx_f32_dit_block_tiny(x, n_tokens, D, d_ff, gamma, eps, Wq, Wk, Wv, Wo, sc1, sh1, gt1, W1, W3, W2, sc2, sh2, gt2, out)
162 if vb != 0 { return 30 }
163 var diff: i64 = 0
164 i = 0
165 while i < n_tokens * D { if out[i] != x[i] { diff = 1 } i = i + 1 }
166 if diff == 0 { return 40 }
167
168 return 0
169}