nx_f32_linear_simd.nx source
↩ module page · 130 lines · 4393 B
1// nx_f32_linear_simd.nx -- SIMD f32 linear (__f32x8_dot) for the SANA-DiT projections, vs scalar: verify+speed.
2//
3// The DiT block's compute bulk is dbl_linear: out[t,o] = Σ_i inp[t,i]*W[o,i]. With packed 4-byte f32, each
4// output is a chunked __f32x8_dot over the in-dim -> ~10x. Verified bit-close vs the i64-f32 scalar linear +
5// timed. This is the f32-SIMD accelerator for the SANA-DiT block's Q/K/V/O/FFN projections.
6// license_tier: ORIGINAL
7import "nx_syscalls.nx"
8import "nx_tier.nx"
9import "nx_le.nx"
10import "nx_strconv.nx"
11import "nx_f32.nx"
12import "nx_f32_div.nx"
13import "nx_f32_cvt.nx"
14import "nx_clock.nx"
15
16// packed-f32 SIMD linear: inp_p [n*id] 4-byte f32, W_p [od*id] 4-byte f32 -> out [n*od] i64-f32
17func dbl_linear_simd(inp_p: *u8, W_p: *u8, out: *i64, n: i64, id: i64, od: i64) -> i64 {
18 var t: i64 = 0
19 while t < n {
20 var o: i64 = 0
21 while o < od {
22 var acc: i64 = 0
23 var c: i64 = 0
24 let ib: i64 = (inp_p as i64) + t * id * 4
25 let wb: i64 = (W_p as i64) + o * id * 4
26 while c < id / 8 {
27 acc = nx_f32_add(acc, __f32x8_dot((ib + c * 32) as *i64, (wb + c * 32) as *i64))
28 c = c + 1
29 }
30 out[t * od + o] = acc
31 o = o + 1
32 }
33 t = t + 1
34 }
35 return 0
36}
37
38// scalar i64-f32 linear reference
39func dbl_linear_scalar(inp: *i64, W: *i64, out: *i64, n: i64, id: i64, od: i64) -> i64 {
40 var t: i64 = 0
41 while t < n {
42 var o: i64 = 0
43 while o < od {
44 var acc: i64 = 0
45 var i: i64 = 0
46 while i < id { acc = nx_f32_add(acc, nx_f32_mul(inp[t * id + i], W[o * id + i])); i = i + 1 }
47 out[t * od + o] = acc
48 o = o + 1
49 }
50 t = t + 1
51 }
52 return 0
53}
54
55func ls_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 {
56 let line: *u8 = sys_mmap(64)
57 var lo: i64 = 0
58 var i: i64 = 0
59 while i < kl { line[lo] = key[i]; lo = lo + 1; i = i + 1 }
60 line[lo] = 0x3D; lo = lo + 1
61 let dec: *u8 = sys_mmap(32)
62 let nd: i64 = nx_strconv_format_i64(v, dec)
63 var k: i64 = 0
64 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 }
65 line[lo] = 0x0A; lo = lo + 1
66 return sys_write(fd, line, lo)
67}
68
69func main() -> i64 {
70 let n: i64 = 8
71 let id: i64 = 512
72 let od: i64 = 64
73 let inp_p: *u8 = sys_mmap(n * id * 4)
74 let W_p: *u8 = sys_mmap(od * id * 4)
75 let inp: *i64 = sys_mmap(n * id * 8) as *i64
76 let W: *i64 = sys_mmap(od * id * 8) as *i64
77 let seven: i64 = nx_i32_to_f32(7)
78 let five: i64 = nx_i32_to_f32(5)
79 var i: i64 = 0
80 while i < n * id {
81 let v: i64 = nx_f32_div(nx_i32_to_f32((i - (i / 7) * 7) + 1), seven)
82 inp[i] = v
83 nx_le_write_u32(inp_p, i * 4, v)
84 i = i + 1
85 }
86 i = 0
87 while i < od * id {
88 let v: i64 = nx_f32_div(nx_i32_to_f32((i - (i / 5) * 5) + 1), five)
89 W[i] = v
90 nx_le_write_u32(W_p, i * 4, v)
91 i = i + 1
92 }
93 let outS: *i64 = sys_mmap(n * od * 8) as *i64
94 let outC: *i64 = sys_mmap(n * od * 8) as *i64
95
96 dbl_linear_simd(inp_p, W_p, outS, n, id, od)
97 dbl_linear_scalar(inp, W, outC, n, id, od)
98 // verify each output within 1%
99 i = 0
100 while i < n * od {
101 let e: i64 = outC[i] & 0x7FFFFFFF
102 let tol: i64 = nx_f32_mul(nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(100)), e)
103 if (nx_f32_sub(outS[i], outC[i]) & 0x7FFFFFFF) >= tol { return 80 }
104 i = i + 1
105 }
106
107 let IT: i64 = 200
108 let t0: i64 = nx_clock_monotonic_ns()
109 var k: i64 = 0
110 while k < IT { dbl_linear_scalar(inp, W, outC, n, id, od); k = k + 1 }
111 let t1: i64 = nx_clock_monotonic_ns()
112 let scalar_ns: i64 = t1 - t0
113 let t2: i64 = nx_clock_monotonic_ns()
114 k = 0
115 while k < IT { dbl_linear_simd(inp_p, W_p, outS, n, id, od); k = k + 1 }
116 let t3: i64 = nx_clock_monotonic_ns()
117 let simd_ns: i64 = t3 - t2
118
119 let ofd: i64 = sys_openat_wr("/tmp/lin_simd.txt" as *u8, 0x1a4)
120 if ofd >= 0 {
121 ls_emit(ofd, "n_out" as *u8, 5, n * od)
122 ls_emit(ofd, "in_dim" as *u8, 6, id)
123 ls_emit(ofd, "scalar_ns" as *u8, 9, scalar_ns)
124 ls_emit(ofd, "simd_ns" as *u8, 7, simd_ns)
125 if simd_ns > 0 { ls_emit(ofd, "speedup_x100" as *u8, 12, scalar_ns * 100 / simd_ns) }
126 sys_close(ofd)
127 }
128 if simd_ns >= scalar_ns { return 81 }
129 return 0
130}