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}