code wiki / (root) / nx_f32_conv2d_backward_fast.nx

nx_f32_conv2d_backward_fast.nx source

↩ module page · 109 lines · 4721 B

1// nx_f32_conv2d_backward_fast.nx -- FAST conv2d backward via im2col + the fork-parallel nx_f32_matmul_t, the 2// companion to nx_f32_conv2d_fast so a full TRAINING STEP is fast (the serial backward is the bottleneck ~2/3 of a 3// pose-net step). Formulation (N=1): colT[K,M] = transposed im2col of input; dWeight[C_out,K] = matmul_t(dOut[C_out,M], 4// colT[K,M]); dCol[M,K] = matmul_t(doutT[M,C_out], weightT[K,C_out]); dInput = col2im(dCol); dBias[co]=sum_p dOut. 5// Bit-exact vs nx_f32_conv2d_backward on integer values. ⚠N=1 (dWeight is assigned by matmul_t, not accumulated). 6// license_tier: ORIGINAL 7import "nx_syscalls.nx" 8import "nx_f32.nx" 9import "nx_f32_matmul_t.nx" 10 11const NX_CVBF_OK: i64 = 0 12const NX_CVBF_ERR_ARGS: i64 = 4 13 14func nx_f32_conv2d_backward_fast(input: *i64, N: i64, C_in: i64, H: i64, W: i64, 15 weight: *i64, C_out: i64, KH: i64, KW: i64, 16 stride: i64, pad: i64, 17 dout: *i64, dinput: *i64, dweight: *i64, dbias: *i64) -> i64 { 18 if stride <= 0 { return NX_CVBF_ERR_ARGS } 19 if N != 1 { return NX_CVBF_ERR_ARGS } 20 let OH: i64 = (H + 2 * pad - KH) / stride + 1 21 let OW: i64 = (W + 2 * pad - KW) / stride + 1 22 if OH <= 0 { return NX_CVBF_ERR_ARGS } 23 if OW <= 0 { return NX_CVBF_ERR_ARGS } 24 let K: i64 = C_in * KH * KW 25 let M: i64 = OH * OW 26 let colT: *i64 = sys_mmap(8 * K * M) as *i64 27 let doutT: *i64 = sys_mmap(8 * M * C_out) as *i64 28 let weightT: *i64 = sys_mmap(8 * K * C_out) as *i64 29 let dcol: *i64 = sys_mmap(8 * M * K) as *i64 30 31 // weightT[k, co] = weight[co, k] 32 var co: i64 = 0 33 while co < C_out { var k: i64 = 0; while k < K { weightT[k * C_out + co] = weight[co * K + k]; k = k + 1 } co = co + 1 } 34 // doutT[p, co] = dout[co, p] 35 co = 0 36 while co < C_out { var p: i64 = 0; while p < M { doutT[p * C_out + co] = dout[co * M + p]; p = p + 1 } co = co + 1 } 37 // colT[k, p] = im2col of input (transposed) 38 var oh: i64 = 0 39 while oh < OH { 40 var ow: i64 = 0 41 while ow < OW { 42 let p: i64 = oh * OW + ow 43 var ci: i64 = 0 44 while ci < C_in { 45 let icb: i64 = ci * H * W 46 var kh: i64 = 0 47 while kh < KH { 48 let ih: i64 = oh * stride + kh - pad 49 var kw: i64 = 0 50 while kw < KW { 51 let iw: i64 = ow * stride + kw - pad 52 var v: i64 = 0 53 if ih >= 0 { if ih < H { if iw >= 0 { if iw < W { v = input[icb + ih * W + iw] } } } } 54 colT[(ci * KH * KW + kh * KW + kw) * M + p] = v 55 kw = kw + 1 56 } 57 kh = kh + 1 58 } 59 ci = ci + 1 60 } 61 ow = ow + 1 62 } 63 oh = oh + 1 64 } 65 66 // dweight[C_out,K] = matmul_t(dout[C_out,M], colT[K,M]) ; dcol[M,K] = matmul_t(doutT[M,C_out], weightT[K,C_out]) 67 nx_f32_matmul_t(dout, colT, dweight, C_out, M, K) 68 nx_f32_matmul_t(doutT, weightT, dcol, M, C_out, K) 69 70 // dbias[co] = sum_p dout[co,p] 71 co = 0 72 while co < C_out { var s: i64 = 0; var p: i64 = 0; while p < M { s = nx_f32_add(s, dout[co * M + p]); p = p + 1 } dbias[co] = s; co = co + 1 } 73 74 // zero dinput, then col2im: scatter dcol[p,k] into dinput[ci, oh*stride+kh-pad, ow*stride+kw-pad] 75 var z: i64 = 0; let nin: i64 = C_in * H * W 76 while z < nin { dinput[z] = 0; z = z + 1 } 77 oh = 0 78 while oh < OH { 79 var ow: i64 = 0 80 while ow < OW { 81 let p: i64 = oh * OW + ow 82 var ci: i64 = 0 83 while ci < C_in { 84 let icb: i64 = ci * H * W 85 var kh: i64 = 0 86 while kh < KH { 87 let ih: i64 = oh * stride + kh - pad 88 if ih >= 0 { if ih < H { 89 var kw: i64 = 0 90 while kw < KW { 91 let iw: i64 = ow * stride + kw - pad 92 if iw >= 0 { if iw < W { 93 let o: i64 = icb + ih * W + iw 94 dinput[o] = nx_f32_add(dinput[o], dcol[p * K + ci * KH * KW + kh * KW + kw]) 95 } } 96 kw = kw + 1 97 } 98 } } 99 kh = kh + 1 100 } 101 ci = ci + 1 102 } 103 ow = ow + 1 104 } 105 oh = oh + 1 106 } 107 sys_munmap(colT as *u8, 8*K*M); sys_munmap(doutT as *u8, 8*M*C_out); sys_munmap(weightT as *u8, 8*K*C_out); sys_munmap(dcol as *u8, 8*M*K) 108 return NX_CVBF_OK 109}