code wiki / (root) / nx_vit_speed_probe.nx

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}