code wiki / (root) / nx_f32_conv2d_fast_gate.nx

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}