conv2d之backward

鱿鱼圈 Lv4

1 填充

在此之前,先补充一些卷积层前向过程中的一个步骤:填充(padding);首先引出问题,例如现在有一个卷积运算具体参数如下: img

卷积的输入 X 宽高都是28,卷积核 W 的宽高是5,取步长 S 为1,那么输出 A 宽高就是24

因为卷积运输的计算方式决定了输出的长度总是小于输入(长度为1的卷积核除外),这样有时会碰到一些问题,例如输入边缘的信息很重要,但是卷积计算对输入的边缘采样不够(计算次数少);或者出于某些原因就是想让卷积前后的数据长度一致

这时候可以原始输入的边缘填充一层‘0’(或几层),使卷积的输入宽高变得更大,然后再进行卷积:

img

如上图,通过给四边各填充2层‘0’,输入的宽高变成了32,卷积后的宽高与原始的输入保持一致

对于填充,常用的有‘SAME’方式和‘VALID’方式,前者是根据卷积核的大小来决定填充的层数,使卷积前后数据大小一致(SAME),填充层数等于卷积核长度整除2;而‘VALID’方式就是不填充;最后,贴一下卷积前后数据宽高的计算公式:

是输出高度, 是输入高度, 是卷积核高度, 是上下填充层数,步长为 ;宽度的计算同理:

2 求输入特征的梯度

2.1 计算公式

现在来讨论卷积层的梯度传递,先贴结果

函数 p(⋅) 作用是对输入进行填充, rot180⁡(W) 是将卷积核旋转180°, 卷积运算符

举个最简单的栗子,如果前向过程中卷积计算情况如下:

img

那么反向过程中梯度传递的情况则是下面这样:

img

从上图中可以看到,将上一层传进来的梯度用0进行了填充,然后取步长为1与卷积核进行卷积运算,就出了本层需要传递的梯度。为了区(fang)分(bian),我将梯度用方块表示,同时卷积核的字体标成了黄色,主要是为了引起注意:卷积核旋转了180°

2.2 公式推导

那么,这个结论是怎么来的?

先看 forward 中的卷积公式(输出位置记为 ):

对某个输入元素 求偏导,由链式法则:

其中,只有当下标对应关系

同时成立时, 才会出现在 的求和项里,因此:

代入即得:

可以看出, 维上随着 增大而取更“靠前”的下标,等价于与旋转后的卷积核做相关运算。

以下面这个简化二维例子说明(单 batch、单通道、单卷积核,):

img

其中 表示对 在四周补零后的取值。

若进一步令:

  • 四周各填充 层 0,得到
  • 顺时针旋转 180°,记为

则在步长为 1 时,可写成与式(1)完全同构的形式:

对比式(1)与式(2):都是“核系数 × 另一张量对应位置”再求和,只是式(1)中输出在 、输入在 ,式(2)在补零后步长为 1,输出在 、梯度在 。因此:

现在,有个两个问题

  • 一个是上面的公式中(下标)越界的处理
  • 另一个时卷积计算步长大于1的处理

第一个容易解决,规定越界为0即可,在编程实现中就是用0对dL/dA进行填充,填充的层数为卷积核大小减1

至于具体的栗子再把前面的图贴一下:

img

至于第二个问题,先说答案,卷积计算中步长大于1的话,则在计算梯度时,先将 dL/dA 用0填充到步长为1的大小,然后再按卷积为1的情况进行处理,下面看具体的栗子:

img

步长为2的卷积

上图是步长为2的卷积运算的栗子,结果就等于在卷积为1的结果之上,按照步长大小间隔取值

所以反向过程计算梯度时就需要先将 dL/dA处理成以下形式:

img

2.3 代码实现

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
// [b, h, w, c] -> [b, out_h * out_w, kh * kw * C]
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;
}


// [D, kh, kw, C] -> [kh * kw * C, D]
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;
}

// [out_h * out_w, kh * kw * C] dot [kh * kw * C, D]->[out_h * out_w, D]
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;
}

// [B, H, W, C] conv [D, KH, KW, C] -> [B, out_h, out_w, D]
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))));
// cout << "维度: [" << B << ' ' << out_h << ' ' << out_w << ' ' << D << "]\n";
for (int b = 0; b < B; b ++) {
vec2d lhs = input_3d[b]; // [out_h * out_w, kh * kw * C]
vec2d rhs = weight_2d; // [kh * kw * C, D]

vec2d res = gemm(lhs, rhs); // [out_h * out_w, D]

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;
}

// [B, out_h, out_w, D] -> [B, H_upsampled, W_upsampled, D]
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;
}

// [B, H_upsampled, W_upsampled, D] -> [B, H_padded, W_padded, D]
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;
}

//[D, KH, KW, C] -> [D, KH, KW, C]
//顺时针旋转180°
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;
}

// [D, KH, KW, C] -> [C, KH, KW, D]
// why?因为套用的conv得模板,输入维度得按固定得顺序,这里需要变化一下维度
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;
}

3 计算卷积核的梯度

3.1 计算公式

现在来讨论卷积层的梯度传递,先贴结果

,⊙卷积运算符

3.2 公式推导

那么,这个结论是怎么来的?

仍从前向公式出发:

注意 只出现在输出通道为 的式子里,因此对某个核元素 求偏导:

在式(1)的求和中,只有当下标满足

时, 的系数才会对 产生贡献,因此:

代入得:

对比式(1)与式(2):式(1)是“核 × 输入”对输出位置求和,式(2)是“上游梯度 × 输入”对 batch 和输出位置求和;二者结构同构,只是把 换成了 ,并把输出下标 保留在求和中。因此:

3.3 代码实现

参考2.3

  • 标题: conv2d之backward
  • 作者: 鱿鱼圈
  • 创建于 : 2025-06-07 22:13:32
  • 更新于 : 2026-06-14 23:33:01
  • 链接: https://yuyanqi.com/2025/06/07/conv2d之backward/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论