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}