nx_vit_speed_probe.nx source
↩ module page · 41 lines · 2203 B
1// nx_vit_speed_probe.nx -- measure the software-f32 matmul rate (MAC/s) to decide if a FAITHFUL f32 ViTPose forward
2// is viable or if the port must ride the int8 path. Times one GEMM A[32,512]@B[512,512] (8.4M MACs) and extrapolates
3// to a full ViTPose-base forward (~17.5 GMAC: patch-embed + 12x[QKV+attn+dense+fc1+fc2] + head conv). expect_exit: 0
4import "nx_syscalls.nx"
5import "nx_f32_cvt.nx"
6import "nx_f32_matmul.nx"
7const K_MAGIC_1000000: i64 = 1000000
8const K_MAGIC_17500000000: i64 = 17500000000
9
10func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 }
11func wn(v: i64) -> i64 { var m: i64=v; if m<0{w("-" as *u8);m=0-m} let t:*u8=sys_mmap(24); var k:i64=0; if m==0{t[0]=48 as u8;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i:i64=0; let o:*u8=sys_mmap(24); while i<k{o[i]=t[k-1-i];i=i+1} sys_write(1,o,k); return 0 }
12
13func main() -> i64 {
14 let M: i64 = 32; let K: i64 = 512; let N: i64 = 512
15 let A: *i64 = sys_mmap(8 * M * K) as *i64
16 let B: *i64 = sys_mmap(8 * K * N) as *i64
17 let C: *i64 = sys_mmap(8 * M * N) as *i64
18 var i: i64 = 0
19 while i < M*K { A[i] = nx_i32_to_f32((i % 5) + 1); i = i + 1 }
20 i = 0
21 while i < K*N { B[i] = nx_i32_to_f32((i % 3) + 1); i = i + 1 }
22
23 let t0: i64 = sys_now_us()
24 nx_f32_matmul(A, B, C, M, K, N)
25 let t1: i64 = sys_now_us()
26
27 let us: i64 = t1 - t0
28 let macs: i64 = M * K * N
29 w("GEMM " as *u8); wn(M); w("x" as *u8); wn(K); w("x" as *u8); wn(N); w(" = " as *u8); wn(macs); w(" MACs in " as *u8); wn(us); w(" us\n" as *u8)
30 if us <= 0 { w("timer resolution too coarse\n" as *u8); sys_exit(0); return 0 }
31 let mac_per_s: i64 = (macs * K_MAGIC_1000000) / us
32 w(" rate = " as *u8); wn(mac_per_s); w(" MAC/s (" as *u8); wn(mac_per_s / K_MAGIC_1000000); w(" MMAC/s)\n" as *u8)
33
34 // full ViTPose-base forward ~= 17.5 GMAC
35 let fwd_macs: i64 = K_MAGIC_17500000000
36 let fwd_us: i64 = (fwd_macs / macs) * us // scale the measured time
37 w(" est full ViTPose forward (~17.5 GMAC) = " as *u8); wn(fwd_us / 1000); w(" ms = " as *u8); wn(fwd_us / K_MAGIC_1000000); w(" s\n" as *u8)
38 w("SPEED-PROBE done\n" as *u8)
39 sys_exit(0)
40 return 0
41}