code wiki / (root) / nx_f32_dit_block_simd.nx

nx_f32_dit_block_simd.nx source

↩ module page · 207 lines · 7872 B

1// nx_f32_dit_block_simd.nx -- SANA-DiT block with SIMD (__f32x8_dot) projections: the end-to-end fast unit. 2// 3// Bridges the two layouts: activations flow as i64-f32 for the elementwise/attention ops; before each 4// projection they PACK to 4-byte f32 (weights pre-packed once) so dbl_linear_simd runs the compute bulk on 5// AVX vmulps (~10x, measured in nx_f32_linear_simd). Same block as nx_f32_dit_block_linear; only the linears 6// change. adaLN-Zero gate (gt=0 => out==x) proves the wiring survives the layout bridging. Scratch mmap'd 7// inside (block called a handful of times, not a hot loop) to keep the arg list small. 8// license_tier: ORIGINAL 9import "nx_syscalls.nx" 10import "nx_le.nx" 11import "nx_f32.nx" 12import "nx_f32_div.nx" 13import "nx_f32_cvt.nx" 14import "nx_f32_activations.nx" 15 16func dsb_phi(v: i64) -> i64 { 17 var r: i64 = v 18 if (v & 0x80000000) != 0 { r = 0 } 19 return nx_f32_add(r, nx_i32_to_f32(1)) 20} 21 22func dsb_pack(src: *i64, dst: *u8, count: i64) -> i64 { 23 var k: i64 = 0 24 while k < count { nx_le_write_u32(dst, k * 4, src[k]); k = k + 1 } 25 return 0 26} 27 28func dsb_linear_simd(inp_p: *u8, W_p: *u8, out: *i64, n: i64, id: i64, od: i64) -> i64 { 29 var t: i64 = 0 30 while t < n { 31 var o: i64 = 0 32 while o < od { 33 var acc: i64 = 0 34 var c: i64 = 0 35 let ib: i64 = (inp_p as i64) + t * id * 4 36 let wb: i64 = (W_p as i64) + o * id * 4 37 while c < id / 8 { acc = nx_f32_add(acc, __f32x8_dot((ib + c * 32) as *i64, (wb + c * 32) as *i64)); c = c + 1 } 38 out[t * od + o] = acc 39 o = o + 1 40 } 41 t = t + 1 42 } 43 return 0 44} 45 46func dsb_modulate(x: *i64, sc: *i64, sh: *i64, out: *i64, n: i64, D: i64) -> i64 { 47 var t: i64 = 0 48 while t < n { 49 var d: i64 = 0 50 while d < D { 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]); d = d + 1 } 51 t = t + 1 52 } 53 return 0 54} 55 56func dsb_linattn(Q: *i64, K: *i64, V: *i64, n: i64, D: i64, out: *i64, S: *i64, z: *i64) -> i64 { 57 var a: i64 = 0 58 while a < D * D { S[a] = 0; a = a + 1 } 59 a = 0 60 while a < D { z[a] = 0; a = a + 1 } 61 var j: i64 = 0 62 while j < n { 63 a = 0 64 while a < D { 65 let pk: i64 = dsb_phi(K[j * D + a]) 66 z[a] = nx_f32_add(z[a], pk) 67 var b: i64 = 0 68 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 } 69 a = a + 1 70 } 71 j = j + 1 72 } 73 var i: i64 = 0 74 while i < n { 75 var den: i64 = 0 76 a = 0 77 while a < D { den = nx_f32_add(den, nx_f32_mul(dsb_phi(Q[i * D + a]), z[a])); a = a + 1 } 78 var b: i64 = 0 79 while b < D { 80 var num: i64 = 0 81 a = 0 82 while a < D { num = nx_f32_add(num, nx_f32_mul(dsb_phi(Q[i * D + a]), S[a * D + b])); a = a + 1 } 83 out[i * D + b] = nx_f32_div(num, den) 84 b = b + 1 85 } 86 i = i + 1 87 } 88 return 0 89} 90 91func dsb_gate_res(x: *i64, y: *i64, gt: *i64, out: *i64, n: i64, D: i64) -> i64 { 92 var t: i64 = 0 93 while t < n { 94 var d: i64 = 0 95 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 } 96 t = t + 1 97 } 98 return 0 99} 100 101func nx_f32_dit_block_simd(x: *i64, n: i64, D: i64, dff: i64, 102 Wq_p: *u8, Wk_p: *u8, Wv_p: *u8, Wo_p: *u8, sc1: *i64, sh1: *i64, gt1: *i64, 103 Wg_p: *u8, Wu_p: *u8, Wd_p: *u8, sc2: *i64, sh2: *i64, gt2: *i64, out: *i64) -> i64 { 104 let hn: *i64 = sys_mmap(n * D * 8) as *i64 105 let hn_p: *u8 = sys_mmap(n * D * 4) 106 let Q: *i64 = sys_mmap(n * D * 8) as *i64 107 let K: *i64 = sys_mmap(n * D * 8) as *i64 108 let V: *i64 = sys_mmap(n * D * 8) as *i64 109 let attn: *i64 = sys_mmap(n * D * 8) as *i64 110 let attn_p: *u8 = sys_mmap(n * D * 4) 111 let ao: *i64 = sys_mmap(n * D * 8) as *i64 112 let x1: *i64 = sys_mmap(n * D * 8) as *i64 113 let S: *i64 = sys_mmap(D * D * 8) as *i64 114 let z: *i64 = sys_mmap(D * 8) as *i64 115 116 dsb_modulate(x, sc1, sh1, hn, n, D) 117 dsb_pack(hn, hn_p, n * D) 118 dsb_linear_simd(hn_p, Wq_p, Q, n, D, D) 119 dsb_linear_simd(hn_p, Wk_p, K, n, D, D) 120 dsb_linear_simd(hn_p, Wv_p, V, n, D, D) 121 dsb_linattn(Q, K, V, n, D, attn, S, z) 122 dsb_pack(attn, attn_p, n * D) 123 dsb_linear_simd(attn_p, Wo_p, ao, n, D, D) 124 dsb_gate_res(x, ao, gt1, x1, n, D) 125 126 let hn2: *i64 = sys_mmap(n * D * 8) as *i64 127 let hn2_p: *u8 = sys_mmap(n * D * 4) 128 let hg: *i64 = sys_mmap(n * dff * 8) as *i64 129 let hu: *i64 = sys_mmap(n * dff * 8) as *i64 130 let hm: *i64 = sys_mmap(n * dff * 8) as *i64 131 let hm_p: *u8 = sys_mmap(n * dff * 4) 132 let ffn: *i64 = sys_mmap(n * D * 8) as *i64 133 dsb_modulate(x1, sc2, sh2, hn2, n, D) 134 dsb_pack(hn2, hn2_p, n * D) 135 dsb_linear_simd(hn2_p, Wg_p, hg, n, D, dff) 136 dsb_linear_simd(hn2_p, Wu_p, hu, n, D, dff) 137 var t: i64 = 0 138 while t < n { 139 var f: i64 = 0 140 while f < dff { hm[t * dff + f] = nx_f32_mul(nx_f32_silu(hg[t * dff + f]), hu[t * dff + f]); f = f + 1 } 141 t = t + 1 142 } 143 dsb_pack(hm, hm_p, n * dff) 144 dsb_linear_simd(hm_p, Wd_p, ffn, n, dff, D) 145 dsb_gate_res(x1, ffn, gt2, out, n, D) 146 return 0 147} 148 149func main() -> i64 { 150 let n: i64 = 16 151 let D: i64 = 64 152 let dff: i64 = 128 153 let x: *i64 = sys_mmap(n * D * 8) as *i64 154 let out: *i64 = sys_mmap(n * D * 8) as *i64 155 let out2: *i64 = sys_mmap(n * D * 8) as *i64 156 let Wq: *i64 = sys_mmap(D * D * 8) as *i64 157 let Wk: *i64 = sys_mmap(D * D * 8) as *i64 158 let Wv: *i64 = sys_mmap(D * D * 8) as *i64 159 let Wo: *i64 = sys_mmap(D * D * 8) as *i64 160 let Wg: *i64 = sys_mmap(dff * D * 8) as *i64 161 let Wu: *i64 = sys_mmap(dff * D * 8) as *i64 162 let Wd: *i64 = sys_mmap(D * dff * 8) as *i64 163 let Wq_p: *u8 = sys_mmap(D * D * 4) 164 let Wk_p: *u8 = sys_mmap(D * D * 4) 165 let Wv_p: *u8 = sys_mmap(D * D * 4) 166 let Wo_p: *u8 = sys_mmap(D * D * 4) 167 let Wg_p: *u8 = sys_mmap(dff * D * 4) 168 let Wu_p: *u8 = sys_mmap(dff * D * 4) 169 let Wd_p: *u8 = sys_mmap(D * dff * 4) 170 let sc1: *i64 = sys_mmap(D * 8) as *i64 171 let sh1: *i64 = sys_mmap(D * 8) as *i64 172 let gt1: *i64 = sys_mmap(D * 8) as *i64 173 let sc2: *i64 = sys_mmap(D * 8) as *i64 174 let sh2: *i64 = sys_mmap(D * 8) as *i64 175 let gt2: *i64 = sys_mmap(D * 8) as *i64 176 let p1: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(10)) 177 let p03: i64 = nx_f32_div(nx_i32_to_f32(3), nx_i32_to_f32(100)) 178 179 var i: i64 = 0 180 while i < n * D { x[i] = nx_f32_div(nx_i32_to_f32((i - (i / 9) * 9) + 1), nx_i32_to_f32(9)); i = i + 1 } 181 i = 0 182 while i < D * D { Wq[i] = p03; Wk[i] = p03; Wv[i] = p03; Wo[i] = p03; i = i + 1 } 183 i = 0 184 while i < dff * D { Wg[i] = p03; Wu[i] = p03; Wd[i] = p03; i = i + 1 } 185 dsb_pack(Wq, Wq_p, D * D); dsb_pack(Wk, Wk_p, D * D); dsb_pack(Wv, Wv_p, D * D); dsb_pack(Wo, Wo_p, D * D) 186 dsb_pack(Wg, Wg_p, dff * D); dsb_pack(Wu, Wu_p, dff * D); dsb_pack(Wd, Wd_p, D * dff) 187 i = 0 188 while i < D { sc1[i] = p1; sh1[i] = p1; sc2[i] = p1; sh2[i] = p1; gt1[i] = 0; gt2[i] = 0; i = i + 1 } 189 190 nx_f32_dit_block_simd(x, n, D, dff, Wq_p, Wk_p, Wv_p, Wo_p, sc1, sh1, gt1, Wg_p, Wu_p, Wd_p, sc2, sh2, gt2, out) 191 i = 0 192 while i < n * D { if out[i] != x[i] { return 10 } i = i + 1 } 193 194 i = 0 195 while i < D { gt1[i] = p1; gt2[i] = p1; i = i + 1 } 196 nx_f32_dit_block_simd(x, n, D, dff, Wq_p, Wk_p, Wv_p, Wo_p, sc1, sh1, gt1, Wg_p, Wu_p, Wd_p, sc2, sh2, gt2, out) 197 nx_f32_dit_block_simd(x, n, D, dff, Wq_p, Wk_p, Wv_p, Wo_p, sc1, sh1, gt1, Wg_p, Wu_p, Wd_p, sc2, sh2, gt2, out2) 198 var diff: i64 = 0 199 i = 0 200 while i < n * D { 201 if out[i] != out2[i] { return 20 } 202 if out[i] != x[i] { diff = 1 } 203 i = i + 1 204 } 205 if diff == 0 { return 21 } 206 return 0 207}