code wiki / (root) / nx_f32_dit_block_tiny.nx

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}