code wiki / (root) / nx_f32_conv2d_backward_fast_gate.nx

nx_f32_conv2d_backward_fast_gate.nx source

↩ module page · 55 lines · 4136 B

1// nx_f32_conv2d_backward_fast_gate.nx -- proves nx_f32_conv2d_backward_fast == nx_f32_conv2d_backward (bit-exact on 2// integer values, stride 1 and 2) for dWeight, dBias, dInput; + measures the speedup. expect_exit: 0 3import "nx_syscalls.nx" 4import "nx_f32_cvt.nx" 5import "nx_f32_conv2d_backward.nx" 6import "nx_f32_conv2d_backward_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) } 10func 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 } 11 12func chk(inp: *i64, W: *i64, dout: *i64, Cin: i64, H: i64, Wd: i64, Cout: i64, KH: i64, KW: i64, st: i64, pd: i64, tag: *u8) -> i64 { 13 let OH: i64=(H+2*pd-KH)/st+1; let OW: i64=(Wd+2*pd-KW)/st+1 14 let dinS: *i64=sys_mmap(8*Cin*H*Wd) as *i64; let dwS: *i64=sys_mmap(8*Cout*Cin*KH*KW) as *i64; let dbS: *i64=sys_mmap(8*Cout) as *i64 15 let dinF: *i64=sys_mmap(8*Cin*H*Wd) as *i64; let dwF: *i64=sys_mmap(8*Cout*Cin*KH*KW) as *i64; let dbF: *i64=sys_mmap(8*Cout) as *i64 16 nx_f32_conv2d_backward(inp,1,Cin,H,Wd, W,Cout,KH,KW, st,pd, dout, dinS, dwS, dbS) 17 nx_f32_conv2d_backward_fast(inp,1,Cin,H,Wd, W,Cout,KH,KW, st,pd, dout, dinF, dwF, dbF) 18 gp(tag) 19 if same(dwS,dwF,Cout*Cin*KH*KW)==1 { if same(dbS,dbF,Cout)==1 { if same(dinS,dinF,Cin*H*Wd)==1 { gp(" dW+dB+dIn bit-exact OK\n" as *u8); return 1 } } } 20 gp(" FAIL\n" as *u8); return 0 21} 22 23func main(argc: i64, argv: *i64) -> i64 { 24 var pass: i64 = 0 25 // small conv C_in=2 C_out=3 4x4 3x3 ; dOut sized per stride 26 let inp: *i64=sys_mmap(8*2*4*4) as *i64; let W: *i64=sys_mmap(8*3*2*3*3) as *i64 27 var i: i64=0; while i<2*4*4 { inp[i]=nx_i32_to_f32((i%7)-3); i=i+1 } 28 i=0; while i<3*2*3*3 { W[i]=nx_i32_to_f32((i%5)-2); i=i+1 } 29 // BF1 stride 1: OH=OW=4 -> dout [3,4,4]=48 30 let d1: *i64=sys_mmap(8*3*4*4) as *i64; i=0; while i<3*4*4 { d1[i]=nx_i32_to_f32((i%6)-2); i=i+1 } 31 if chk(inp,W,d1, 2,4,4, 3,3,3, 1,1, "BF1 stride-1:" as *u8)==1 { pass=pass+1 } 32 // BF2 stride 2: OH=OW=2 -> dout [3,2,2]=12 33 let d2: *i64=sys_mmap(8*3*2*2) as *i64; i=0; while i<3*2*2 { d2[i]=nx_i32_to_f32((i%6)-2); i=i+1 } 34 if chk(inp,W,d2, 2,4,4, 3,3,3, 2,1, "BF2 stride-2:" as *u8)==1 { pass=pass+1 } 35 36 // BF3 speedup: C_in=16 C_out=48 64x48 3x3 pad1 37 let bin: *i64=sys_mmap(8*16*64*48) as *i64; let bW: *i64=sys_mmap(8*48*16*3*3) as *i64; let bd: *i64=sys_mmap(8*48*64*48) as *i64 38 i=0; while i<16*64*48 { bin[i]=nx_i32_to_f32((i%9)-4); i=i+1 } 39 i=0; while i<48*16*3*3 { bW[i]=nx_i32_to_f32((i%5)-2); i=i+1 } 40 i=0; while i<48*64*48 { bd[i]=nx_i32_to_f32((i%7)-3); i=i+1 } 41 let dinS: *i64=sys_mmap(8*16*64*48) as *i64; let dwS: *i64=sys_mmap(8*48*16*3*3) as *i64; let dbS: *i64=sys_mmap(8*48) as *i64 42 let dinF: *i64=sys_mmap(8*16*64*48) as *i64; let dwF: *i64=sys_mmap(8*48*16*3*3) as *i64; let dbF: *i64=sys_mmap(8*48) as *i64 43 let s0: i64=sys_now_us(); nx_f32_conv2d_backward(bin,1,16,64,48, bW,48,3,3,1,1, bd, dinS, dwS, dbS); let s1: i64=sys_now_us() 44 let f0: i64=sys_now_us(); nx_f32_conv2d_backward_fast(bin,1,16,64,48, bW,48,3,3,1,1, bd, dinF, dwF, dbF); let f1: i64=sys_now_us() 45 let eq: i64=same(dwS,dwF,48*16*3*3)*same(dbS,dbF,48)*same(dinS,dinF,16*64*48) 46 gp("BF3 serial=" as *u8); gn((s1-s0)/1000); gp("ms fast=" as *u8); gn((f1-f0)/1000); gp("ms speedup=" as *u8) 47 if f1-f0>0 { gn(((s1-s0)*100)/(f1-f0)); gp("/100x bit-exact=" as *u8); gn(eq); gp("\n" as *u8) } 48 if eq==1 { if f1-f0 < s1-s0 { pass=pass+1; gp("BF3 fast backward faster + bit-exact OK\n" as *u8) } } 49 if eq==0 { gp("BF3 FAIL not bit-exact\n" as *u8) } 50 51 gp("CONV-BWD-FAST-GATE pass=" as *u8); gn(pass); gp("/3\n" as *u8) 52 if pass==3 { gp("CONV-BWD-FAST-GATE GREEN 3/3 (fast backward == serial, and faster) -- full fast conv for training\n" as *u8); sys_exit(0) } 53 sys_exit(1) 54 return 0 55}