1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185
| vec3d im2col_lhs_new(vec4d v, int kh, int kw, int sh, int sw) {
int B = v.size(), H = v[0].size(), W = v[0][0].size(), C = v[0][0][0].size();
int out_h = (H - kh) / sh + 1; int out_w = (W - kw) / sw + 1;
vec3d output(B, vec2d(out_h * out_w, vec1d(kh * kw * C)));
for (int b = 0; b < B; b ++) { int row = 0; for (int h = 0; h < out_h; h ++) for (int w = 0; w < out_w; w ++) { int idx = 0; for (int i = 0; i < kh; i ++) for (int j = 0; j < kw; j ++) for (int c = 0; c < C; c ++) output[b][row][idx ++] = v[b][sh * h + i][sw * w + j][c]; row ++; } }
return output; }
vec2d im2col_rhs_new(vec4d v) { int Dout = v.size(); int Hout = v[0].size(); int Wout = v[0][0].size(); int Cout = v[0][0][0].size(); vec2d res(Hout * Wout * Cout, vec1d(Dout));
for (int d = 0; d < Dout; d ++) for (int i = 0; i < Hout; i ++) for (int j = 0; j < Wout; j ++) for (int k = 0; k < Cout; k ++) res[i * Wout * Cout + j * Cout + k][d] = v[d][i][j][k]; return res; }
vec2d gemm(vec2d a, vec2d b) { int n = a.size(), m = a[0].size(), l = b[0].size(); vec2d c(n, vec1d(l)); for (int i = 0; i < n; i ++) for (int k = 0; k < m; k ++) for (int j = 0; j < l; j ++) c[i][j] += a[i][k] * b[k][j]; return c; }
vec4d conv2d(vec4d input, vec4d weight, int sh, int sw) {
int H = input[0].size(); int W = input[0][0].size(); int D = weight.size(); int kh = weight[0].size(); int kw = weight[0][0].size(); int out_h = (H - kh) / sh + 1; int out_w = (W - kw) / sw + 1;
vec3d input_3d = im2col_lhs_new(input, kh, kw, sh, sw); vec2d weight_2d = im2col_rhs_new(weight);
int B = input.size();
vec4d ans = vec4d(B, vec3d(out_h, vec2d(out_w, vec1d(D)))); for (int b = 0; b < B; b ++) { vec2d lhs = input_3d[b]; vec2d rhs = weight_2d;
vec2d res = gemm(lhs, rhs);
for (int d = 0; d < D; d ++) for (int i = 0; i < out_h; i ++) for (int j = 0; j < out_w; j ++) ans[b][i][j][d] += res[i * out_w + j][d]; }
return ans; }
vec4d Unsampling(vec4d v, int sh, int sw) { int B = v.size(); int H = v[0].size(); int W = v[0][0].size(); int C = v[0][0][0].size();
int H_upsampled = (H - 1) * sh + 1; int W_upsampled = (W - 1) * sw + 1;
vec4d unsampled(B, vec3d(H_upsampled, vec2d(W_upsampled, vec1d(C, 0))));
for (int b = 0; b < B; b ++) for (int h = 0; h < H; h ++) for (int w = 0; w < W; w ++) { int rel_h = h * sh; int rel_w = w * sw; unsampled[b][rel_h][rel_w] = v[b][h][w]; }
return v; }
vec4d Padding(vec4d v, int kh, int kw) { int B = v.size(); int H_upsampled = v[0].size(); int W_upsampled = v[0][0].size(); int C = v[0][0][0].size(); int pad_top = kh - 1; int pad_left = kw - 1; int H_padded = H_upsampled + 2 * pad_top; int W_padded = W_upsampled + 2 * pad_left;
vec4d padded(B, vec3d(H_padded, vec2d(W_padded, vec1d(C, 0))));
for (int b = 0; b < B; b ++) for (int h = 0; h < H_upsampled; h ++) if (h + pad_top >= 0 && h + pad_top < H_padded) for (int w = 0; w < W_upsampled; w ++) if (w + pad_left >= 0 && w + pad_left < W_padded) padded[b][h + pad_top][w + pad_left] = v[b][h][w];
return padded; }
vec4d kernel_rotate180(vec4d kernel) { int D = kernel.size(); int KH = kernel[0].size(); int KW = kernel[0][0].size(); int C = kernel[0][0][0].size();
vec4d out(D, vec3d(KH, vec2d(KW, vec1d(C))));
for (int d = 0; d < D; d ++) for (int h = 0; h < KH; h ++) for (int w = 0; w < KW; w ++) for (int c = 0; c < C; c ++) out[d][h][w][c] = kernel[d][KH - 1 - h][KW - 1 - w][c];
return out; }
vec4d transpose_d_c(vec4d weight) { int D = weight.size(); int KH = weight[0].size(); int KW = weight[0][0].size(); int C = weight[0][0][0].size();
vec4d out(D, vec3d(KH, vec2d(KW, vec1d(C))));
for (int c = 0; c < C; c ++) for (int h = 0; h < KH; h ++) for (int w = 0; w < KW; w ++) for (int d = 0; d < D; d ++) out[c][h][w][d] = weight[d][h][w][c];
return out; }
vec4d conv2d_backward_dx(vec4d input, vec4d kernel, int sh, int sw) { int kh = kernel[0].size(); int kw = kernel[0][0].size(); auto unsampled_input = Unsampling(input, sh, sw); auto padded_input = Padding(unsampled_input, kh, kw); auto rot_kernel = kernel_rotate180(kernel); auto permuted_kernel = transpose_d_c(rot_kernel); auto ans = conv2d(padded_input, permuted_kernel, sh, sw); return ans; }
|