code wiki / (root) / nx_f32_linear_simd.nx

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}