nx_f32_conv2d_fast_gate.nx source
↩ module page · 56 lines · 3469 B
1// nx_f32_conv2d_fast_gate.nx -- proves nx_f32_conv2d_fast == nx_f32_conv2d_forward (bit-exact on integer values,
2// stride 1 and 2) + measures the im2col/matmul_t speedup on a student-sized conv. expect_exit: 0
3import "nx_syscalls.nx"
4import "nx_f32_cvt.nx"
5import "nx_f32_conv2d.nx"
6import "nx_f32_conv2d_fast.nx"
7
8func gp(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} return sys_write(1,s,n) }
9func gn(v: i64) -> i64 { let bb:*u8=sys_mmap(28); var m:i64=v; if m<0{sys_write(1,"-" as *u8,1);m=0-m} let t:*u8=sys_mmap(28); 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; while i<k{bb[i]=t[k-1-i];i=i+1} return sys_write(1,bb,k) }
10
11// compare two f32 buffers element-wise (bit-exact)
12func same(a: *i64, b: *i64, n: i64) -> i64 { var i: i64=0; while i<n { if a[i]!=b[i] { return 0 } i=i+1 } return 1 }
13
14func main(argc: i64, argv: *i64) -> i64 {
15 var pass: i64 = 0
16
17 // F1 equivalence stride 1: C_in=2, C_out=3, 4x4, 3x3, pad1 -> 4x4 out
18 let inp: *i64=sys_mmap(8*2*4*4) as *i64
19 let W: *i64=sys_mmap(8*3*2*3*3) as *i64
20 let B: *i64=sys_mmap(8*3) as *i64
21 var i: i64=0; while i<2*4*4 { inp[i]=nx_i32_to_f32((i%7)-3); i=i+1 }
22 i=0; while i<3*2*3*3 { W[i]=nx_i32_to_f32((i%5)-2); i=i+1 }
23 B[0]=nx_i32_to_f32(1); B[1]=nx_i32_to_f32(0-2); B[2]=nx_i32_to_f32(3)
24 let os: *i64=sys_mmap(8*3*4*4) as *i64; let of: *i64=sys_mmap(8*3*4*4) as *i64
25 nx_f32_conv2d_forward(inp,1,2,4,4, W,3,3,3,1,1, B, os)
26 nx_f32_conv2d_fast(inp,1,2,4,4, W,3,3,3,1,1, B, of)
27 if same(os, of, 3*4*4)==1 { pass=pass+1; gp("F1 stride-1 fast==serial bit-exact OK\n" as *u8) } else { gp("F1 FAIL\n" as *u8) }
28
29 // F2 equivalence stride 2: 4x4 -> 2x2
30 let os2: *i64=sys_mmap(8*3*2*2) as *i64; let of2: *i64=sys_mmap(8*3*2*2) as *i64
31 nx_f32_conv2d_forward(inp,1,2,4,4, W,3,3,3,2,1, B, os2)
32 nx_f32_conv2d_fast(inp,1,2,4,4, W,3,3,3,2,1, B, of2)
33 if same(os2, of2, 3*2*2)==1 { pass=pass+1; gp("F2 stride-2 fast==serial bit-exact OK\n" as *u8) } else { gp("F2 FAIL\n" as *u8) }
34
35 // F3 speedup: C_in=16, C_out=48, 64x48, 3x3, pad1 (student conv2 size, ~21M MACs)
36 let bin: *i64=sys_mmap(8*16*64*48) as *i64
37 let bW: *i64=sys_mmap(8*48*16*3*3) as *i64
38 let bB: *i64=sys_mmap(8*48) as *i64
39 i=0; while i<16*64*48 { bin[i]=nx_i32_to_f32((i%9)-4); i=i+1 }
40 i=0; while i<48*16*3*3 { bW[i]=nx_i32_to_f32((i%5)-2); i=i+1 }
41 i=0; while i<48 { bB[i]=0; i=i+1 }
42 let bos: *i64=sys_mmap(8*48*64*48) as *i64; let bof: *i64=sys_mmap(8*48*64*48) as *i64
43 let ts0: i64=sys_now_us(); nx_f32_conv2d_forward(bin,1,16,64,48, bW,48,3,3,1,1, bB, bos); let ts1: i64=sys_now_us()
44 let tf0: i64=sys_now_us(); nx_f32_conv2d_fast(bin,1,16,64,48, bW,48,3,3,1,1, bB, bof); let tf1: i64=sys_now_us()
45 let serial_us: i64=ts1-ts0; let fast_us: i64=tf1-tf0
46 let eq: i64=same(bos, bof, 48*64*48)
47 gp("F3 serial=" as *u8); gn(serial_us/1000); gp("ms fast=" as *u8); gn(fast_us/1000); gp("ms speedup=" as *u8)
48 if fast_us>0 { gn((serial_us*100)/fast_us); gp("/100x bit-exact=" as *u8); gn(eq); gp("\n" as *u8) }
49 if eq==1 { if fast_us < serial_us { pass=pass+1; gp("F3 fast is faster AND bit-exact OK\n" as *u8) } }
50 if eq==0 { gp("F3 FAIL not bit-exact\n" as *u8) }
51
52 gp("CONV-FAST-GATE pass=" as *u8); gn(pass); gp("/3\n" as *u8)
53 if pass==3 { gp("CONV-FAST-GATE GREEN 3/3 (im2col+matmul_t == serial, and faster)\n" as *u8); sys_exit(0) }
54 sys_exit(1)
55 return 0
56}