资讯动态

手撕Conv2D:从零实现前向传播、反向传播与im2col优化

发布时间:2026/10/3 13:14:00 来源:尧图企业网站定制
我做面试官这几年几乎每次招算法岗都会让人“手撕代码Conv2D”。为什么这道题这么经典因为它表面上考的是你会不会背卷积公式实际上考的是你对数据流、索引映射、边界条件和反向传播的理解深度。背过 PyTorch 谁都会调但真要你不查文档在白板上把一个二维卷积的前向和反向写对很多人会在 padding、stride 和四层循环上翻车。这篇博文我把自己手撕 Conv2D 的完整过程整理出来不讲虚的直接用代码和推导说话。内容包括最朴素的纯 Python 实现、和 PyTorch 的对拍验证、反向传播的手推公式、以及 im2col 这种工业级优化思路。无论你是准备技术面试、正在读源码还是想加深对 CNN 底层的理解这篇都能给你一套可以直接落地的参考方案。1. 动手前先把卷积的数据映射理清楚1.1 输入输出形状与符号约定手撕卷积最大的障碍不是“不会写循环”而是从头到尾的索引符号不一致。所以我先约定一套清晰的记号后面所有代码都基于它避免一边写一边怀疑人生。假设输入是标准 PyTorch 布局 NCHW即x: (N, C_in, H_in, W_in)weight: (C_out, C_in, K_h, K_w)bias: (C_out,)输出 y: (N, C_out, H_out, W_out)输出尺寸的计算公式是H_out floor((H_in 2 * pad_h - K_h) / stride_h) 1W_out floor((W_in 2 * pad_w - K_w) / stride_w) 1这个公式一定不要记错它是后面所有索引计算的基石。默认情况下卷积核都是正方形即 K_h K_w Kpad 和 stride 也往往在宽高方向取相同值。但是成熟实现里宽高可以不一致所以手撕时可以把它们拆开写这样以后要支持 asymmetric 卷积也不用改结构。完整的前向映射关系如下y[n][oc][oh][ow] bias[oc] Σ(ic0..C_in-1) Σ(kh0..K_h-1) Σ(kw0..K_w-1) x[n][ic][oh * stride_h kh - pad_h][ow * stride_w kw - pad_w] * weight[oc][ic][kh][kw]这里的关键是输入坐标的对应输出位置 (oh, ow) 对应的输入感受野中心是被 stride 拉开的索引是 oh * stride_h kh - pad_h。很多初写者会在这一步把 kh 和 ow 的关系搞反或者忘了减 padding。我建议你在纸上画一个 5x5 输入、3x3 核、stride1、pad1 的图把输出第一个元素和输入哪些位置相关标出来这一步永远不亏。1.2 明确“手撕”的边界边界条件所谓手撕代码一般有三种语境。第一种是白板面试要求在十几分钟内写出可运行的 Conv2D 前向通常不需要真的跑结果第二种是本地实现要求写一个能跑通并且和 PyTorch 数值对齐的函数第三种是源码级实现比如用 C/CUDA 重写操作符。这篇博文针对前两种但我会把第三种需要的思路也带上因为 im2col 和反向传播的手推直接对应框架里的实际做法。在写代码之前还要确定一个事情padding 怎么处理。最直观的做法是先用 np.pad 给输入填充一圈 0再老老实实做卷积。这种做法写起来简单面试时不容易漏边界但缺点是当输入很大时浪费内存也增加了拷贝时间。另一种做法是“边界判断”循环里检查输入坐标是否落在 [0, H_in) 和 [0, W_in) 区间内如果越界就跳过本次累加。这种方式不额外分配内存更接近生产级实现的思路。我下文的朴素版采用边界判断因为这样既能展示对 padding 虚拟坐标的理解又不会引入额外的内存开销。2. 第一版可运行实现最朴素的四重循环2.1 四重循环的骨架代码最直接的实现就是按照公式机械翻译外层遍历 batch 和输出通道内层遍历输出高度和宽度最内层累加输入通道和卷积核的乘积。我用函数 conv2d_forward_naive 来写。import numpy as np def conv2d_forward_naive(x, weight, bias, stride1, pad0): 最朴素的 Conv2D 前向实现。 x: (N, C_in, H_in, W_in) float32 weight: (C_out, C_in, K_h, K_w) float32 bias: (C_out,) float32 N, C_in, H_in, W_in x.shape C_out, _, K_h, K_w weight.shape H_out (H_in 2 * pad - K_h) // stride 1 W_out (W_in 2 * pad - K_w) // stride 1 y np.zeros((N, C_out, H_out, W_out), dtypenp.float32) for n in range(N): for oc in range(C_out): for oh in range(H_out): for ow in range(W_out): acc bias[oc].astype(np.float32) for ic in range(C_in): for kh in range(K_h): for kw in range(K_w): ih oh * stride kh - pad iw ow * stride kw - pad if 0 ih H_in and 0 iw W_in: acc x[n, ic, ih, iw] * weight[oc, ic, kh, kw] y[n, oc, oh, ow] acc return y这段代码的结构非常清晰三层输出索引 三层输入索引总共有六层循环。由于每个输出元素的计算彼此独立这个版本天然支持并行化只不过纯 Python 循环跑起来很慢只能用来验证逻辑或处理小尺寸输入。如果把最内层的 ic/kh/kw 循环用向量化替换速度就能提升不少但那是优化阶段的事第一版先保证绝对正确。2.2 为什么先写朴素版而不是直接上 im2col我知道看到这里有经验的读者可能会问“现代框架根本不是这么干的你教这个干嘛”这个问题的答案恰好是理解 CNN 底层最重要的部分。朴素四重循环是卷积的数学定义它把每个输出像素和输入感受野的对应关系写得明明白白。你只有先写过一遍这种版本才能在调试模型、排查性能瓶颈时感知到卷积的访问模式。比如当你发现某层 GPU 利用率很低时如果心里没有“卷积是滑窗、每个窗口要做 K 乘累加”的底层画面就很难判断到底是访存受限还是计算受限。im2col、Winograd、FFT 这些优化全部是从朴素版出发做等价变换跳过朴素版直接追优化容易迷失方向。另一方面面试场景中朴素版也是最好的“锚点”。面试官问你“怎么优化”你可以在朴素版基础上逐步展开但如果你一上来就写 im2col代码复杂度会直线上升白板根本写不完反而把自己绕进去。所以我的习惯是先给朴素版跑通再给优化思路两段代码配合着讲既展示基础又展示水平。3. 对拍验证和 PyTorch 对齐才算真的对3.1 为什么必须做数值对拍手撕代码最容易犯的错误是“看起来对了实际上差了 padding 一圈”或者“stride 索引从 1 开始而不是从 0 开始”。这类错误靠人眼很难发现尤其是当输入输出尺寸比较整的时候错误值可能被巧合掩盖。所以写完之后一定不能靠感觉判断必须拿 PyTorch 的 nn.Conv2d 作为标准答案让两组结果逐元素做对比。对拍时要注意一个细节PyTorch 的卷积默认开启了 bias如果你的实现偏置维度写错了误差很可能在某个通道整体偏移。所以测试用例要设计成能分辨问题来源的形式最好分别跑带 bias 和不带 bias 两组。另外 dtypes 要一致卷积内部累加用 float32对比时不能把一边转成 float64 另一边保持 float32那样的误差没有参考意义。3.2 对拍测试代码与实测结果我构造一个“非整数对齐”的用例让 padding、stride 和核尺寸互相配合避免因为巧合掩盖问题import torch import torch.nn as nn def test_conv_forward(): torch.manual_seed(0) N, C_in, H, W 2, 3, 9, 11 C_out, K, stride, pad 4, 3, 2, 1 x_np np.random.randn(N, C_in, H, W).astype(np.float32) w_np np.random.randn(C_out, C_in, K, K).astype(np.float32) b_np np.random.randn(C_out).astype(np.float32) # numpy 朴素实现 y_numpy conv2d_forward_naive(x_np, w_np, b_np, stridestride, padpad) # PyTorch 参考实现 x_t torch.from_numpy(x_np) conv nn.Conv2d(C_in, C_out, K, stridestride, paddingpad, biasTrue) conv.weight.data torch.from_numpy(w_np) conv.bias.data torch.from_numpy(b_np) y_torch conv(x_t).detach().numpy() diff np.abs(y_numpy - y_torch) print(max abs diff:, diff.max()) print(mean abs diff:, diff.mean()) test_conv_forward()我实际跑下来的结果是 max abs diff 在 1e-6 到 1e-5 之间。这个量级的差异不是错误而是运算顺序不同导致的浮点舍入噪声。只要最大绝对误差不超过 1e-4并且误差分布随机没有结构性偏移就说明索引映射、累加逻辑和偏置处理都正确。如果误差很大比如超过 1e-2那就不是舍入问题了一定是某个索引关系写错了。这个时候建议缩小到单通道、单 batch、小尺寸的例子手工算一个输出元素出来对比定位速度会快很多。4. 提速关键im2col 与矩阵乘法视角4.1 im2col 的原理把滑窗拉平朴素版虽然正确但速度实在没法用于生产环境。工业界最常用的优化手段之一是 im2col它的核心思想特别简单卷积在数学上是“滑窗 内积”而内积正是矩阵乘法做的事情。所以如果能提前把每个滑窗对应的一小块输入数据拉成一个向量再把所有滑窗的向量拼成一个大矩阵那么卷积就变成了一次矩阵乘法。具体来说输入 x 经过 im2col 后会变成一个二维矩阵 col形状大约是 (C_in * K_h * K_w, N * H_out * W_out)。卷积核 weight 被 reshape 成 (C_out, C_in * K_h * K_w)。两个矩阵相乘后结果正好是 (C_out, N * H_out * W_out)再 reshape 成输出格式即可。这个变换的本质是用空间换时间牺牲内存、换取矩阵乘法的高效执行。如果输入是 224x224、通道 64、核 3x3expand 出来的 col 矩阵体积会非常大所以工业实现通常会分块处理避免一次性撑爆内存。这也是为什么实际框架里很少看到一次把整张图展开而是按 tile 展开、乘完就丢。4.2 一个简单的 im2col 实现示例下面给出一个基于 numpy 的简化版 im2col目的是帮助你理解形状变化的过程而不是造一个生产级轮子def im2col(x, K_h, K_w, stride, pad): N, C_in, H_in, W_in x.shape H_out (H_in 2 * pad - K_h) // stride 1 W_out (W_in 2 * pad - K_w) // stride 1 x_pad np.pad(x, ((0, 0), (0, 0), (pad, pad), (pad, pad)), modeconstant) cols np.zeros((N, C_in, K_h, K_w, H_out, W_out), dtypex.dtype) for i in range(K_h): for j in range(K_w): cols[:, :, i, j, :, :] x_pad[:, :, i:i H_out * stride:stride, j:j W_out * stride:stride] cols cols.transpose(1, 2, 3, 0, 4, 5).reshape(C_in * K_h * K_w, -1) return cols, H_out, W_out # 使用示例 x_np np.random.randn(2, 3, 9, 11).astype(np.float32) cols, H_out, W_out im2col(x_np, 3, 3, stride2, pad1) w_matrix w_np.reshape(w_np.shape[0], -1) # (C_out, C_in*K_h*K_w) out_matrix w_matrix cols # (C_out, N*H_out*W_out)这段代码里最容易出错的地方是步长切片i:i H_out * stride:stride的写法很绕但它的本质是把每种卷积核偏移下的采样点一次性取出来。你可以在小例子上打印 x_pad 和 cols 的形状手动对照一遍输出会非常有帮助。需要说明的是这个实现的内存效率并不高真实框架还会做更多优化比如用共享内存、提前 padding 后按 tile 拷贝等。但作为手撕题能把 im2col 的思路和核心代码讲清楚已经足够展示对卷积计算本质的理解了。4.3 矩阵乘法版本与朴素版的对比把 im2col 得到的矩阵乘出来之后再 reshape 回正常输出。相比朴素版矩阵乘法版本的主要收益是np.dot / BLAS 对矩阵乘法做了高度优化计算密度高、缓存友好即使引入了额外内存拷贝整体速度仍然明显快于六层 Python 循环。实测同样输入下能快一到两个数量级尺寸越大、通道越多收益越明显。但是要记住一个关键点im2col 并不是唯一优化路径。Winograd 在 3x3 卷积上有更少的乘法次数FFT 在大核卷积上另有一套优势Tensor Cores 又走的是低精度矩阵乘融合的路线。手撕题中不需要全部掌握但知道 im2col 的存在和局限会让你在后续看 Kernel 源码时少很多障碍。5. 反向传播手撕中最容易翻车的地方5.1 链式法则下的三组梯度公式前向写完下一关就是反向传播。Conv2D 的反向要计算三个梯度dL/dW、dL/db 和 dL/dx。假设上游传给本层的梯度是 dout形状和 y 相同那么根据链式法则我们可以分三步推导。偏置的梯度最简单。因为 y[n][oc][oh][ow] 的计算里bias[oc] 是一个常数项所以 dout 对 bias[oc] 求导时只需要把所有 batch、oh、ow 位置的梯度加起来db[oc] Σ_n Σ_oh Σ_ow dout[n][oc][oh][ow]权重的梯度稍微复杂一点。权重 weight[oc][ic][kh][kw] 影响了所有输出位置所以梯度是把每个输出位置的梯度乘以对应的输入像素值再累加。公式是dW[oc][ic][kh][kw] Σ_n Σ_oh Σ_ow dout[n][oc][oh][ow] * x[n][ic][oh * stride kh - pad][ow * stride kw - pad]注意这里输入坐标同样要经过 padding 偏移和 stride 缩放边界外的值不参与累加。输入的梯度 dA 是最容易出错的。它的语义是把梯度从输出“散射”回输入。每个输出梯度 dout[n][oc][oh][ow] 都会对输入感受野内所有位置产生贡献。换句话说我们需要遍历该输出位置关联的输入坐标把 dout 乘以对应的 weight 再加回 dA 对应位置。由于一个输入像素会被多个输出滑窗覆盖梯度必须累加而不是覆盖。5.2 反向传播代码实现def conv2d_backward_naive(dout, x, weight, bias, stride1, pad0): N, C_in, H_in, W_in x.shape C_out, _, K_h, K_w weight.shape H_out (H_in 2 * pad - K_h) // stride 1 W_out (W_in 2 * pad - K_w) // stride 1 dx np.zeros_like(x) dw np.zeros_like(weight) db np.zeros_like(bias) for oc in range(C_out): db[oc] np.sum(dout[:, oc, :, :]) for n in range(N): for oc in range(C_out): for oh in range(H_out): for ow in range(W_out): grad_out dout[n, oc, oh, ow] for ic in range(C_in): for kh in range(K_h): for kw in range(K_w): ih oh * stride kh - pad iw ow * stride kw - pad if 0 ih H_in and 0 iw W_in: dw[oc, ic, kh, kw] x[n, ic, ih, iw] * grad_out dx[n, ic, ih, iw] weight[oc, ic, kh, kw] * grad_out return dx, dw, db这段实现把 dW 和 dx 放在同一个循环里算因为访问的输入坐标是相同的。注意 dw 的累加方向每个输入像素在不同的输出窗口中出现了多次所以 dw 是把所有可能位置累加到一起这和卷积核滑窗的“共享权重”直接相关。5.3 反向对拍用 torch.autograd.grad 验证写反向时同样必须对拍。PyTorch 的做法是构造一个简单计算图调用 backward再用 grad 接口取出梯度def test_conv_backward(): torch.manual_seed(1) N, C_in, H, W 2, 3, 8, 10 C_out, K, stride, pad 4, 3, 1, 1 x_t torch.randn(N, C_in, H, W, requires_gradTrue) conv nn.Conv2d(C_in, C_out, K, stridestride, paddingpad, biasTrue) y_t conv(x_t) y_t.sum().backward() dx_t x_t.grad dw_t conv.weight.grad db_t conv.bias.grad x_np x_t.detach().numpy() w_np conv.weight.detach().numpy() b_np conv.bias.detach().numpy() dout_np torch.ones_like(y_t).numpy() dx_np, dw_np, db_np conv2d_backward_naive(dout_np, x_np, w_np, b_np, stridestride, padpad) print(dx max diff:, np.abs(dx_np - dx_t.numpy()).max()) print(dw max diff:, np.abs(dw_np - dw_t.numpy()).max()) print(db max diff:, np.abs(db_np - db_t.numpy()).max())这里 dout 用全 1 的意思是简化验证因为y_t.sum().backward()产生的梯度恰好就是全 1。如果测试采用随机 dout就要在 PyTorch 里手动构造上游梯度做法是y_t.backward(dout_t)。实测跑下来dx、dw、db 的最大误差都在 1e-5 量级这已经足够说明实现正确。反向传播手撕有一个小技巧dw 和 dx 要一起验证因为两者的错误往往成对出现。如果你发现 dw 不对但 dx 对大概率是累加顺序或者 kh/kw 索引错位。6. 手撕 Conv2D 的踩坑速查表与调试方法6.1 高频错误与排查思路我把这几年带新人、刷题、面候选人的过程中反复遇到的错误整理成一个速查表方便你本地写完代码后逐项排查错误类型现象排查方法padding 偏移忘记减输出尺寸对但整体图像错位打印第一个输出行手算一个小例子对拍stride 参与索引时写错输出尺寸不对或越界检查 oh * stride kh - pad单独跑 stride1 看是否正常bias 忘记加或变成逐像素加所有通道整体偏移或梯度完全不对先跑不带 bias 的用例再跑带 bias 的用例对比边界判断写成ih H_in少了等号图片最右侧/最下侧一条梯度或值异常使用 H_in偶数、小核、小 stride 构造单点感受野用例反向 dw 累加顺序颠倒dw 梯度形状对但数值错误用权重的中心元素单独验算dx 没有累加而是覆盖反向梯度比预期小很多或出现规则条纹检查 dx 赋值是否用了浮点误差阈值判断太严格明明正确但 max diff 报 1e-5放宽到 1e-4 以内观察是否有结构性偏差排查时最好的手段是“极小尺寸手工验算”。比如输入 1x1、单通道、3x3 核、stride2、pad1输出尺寸算出来是多少手算一个元素作为基准。这样的用例能瞬间暴露索引映射错误。不要一上来就在大输入上找 bug那样只会淹没在输出矩阵里。6.2 我的几条实操心得第一先写一个帮助函数打印输入形状、输出形状、权重形状确保和预期一致。很多时候代码没问题但测试用例构造反了比如把 weight 的通道序传错结果差异巨大还找不到原因。第二单独验证前向后再写反向。千万不要把前向反向一口气写完再整体调试。因为反向的错误可能来自前向的隐藏偏差也可能来自反向自己的索引错误叠加在一起非常难排查。我在实际写手撕题时严格保持“前向对拍通过 → 反向对拍通过”的顺序这个习惯帮我省了大量时间。第三不要忽略 stride 和 pad 同时存在的测试用例。很多面试题为了图省事只考 stride1、pad0但这会掩盖掉一半以上的边界错误。我自己写代码时至少会跑 stride1, pad0 和 stride2, pad1 两组数据一组验证中心区域一组验证边缘区域两者都对了才敢说实现正确。7. 扩展手撕代码之后还能往哪些方向想如果面试或者平时研究已经做到这一步还可以往更深的方向延展。比如把 Conv2D 改成 depthwise 卷积只需要把 C_out 和 C_in 的对应关系变成组卷积通常组数等于输入通道改成空洞卷积只需要把核索引从 kh, kw 变成 kh * dilation, kw * dilation。这些变体的手感很接近但面试里提前准备好会有加分效果。另外可以想想反向传播里最耗时间的部分在哪。很明显是内层循环遍历 K_h * K_w 且每个位置都要做一次乘加。工业实现里 Winograd 或 FFT 就是为了减少这个内层循环的乘法次数尤其在小核 3x3 场景下Winograd F(2, 3) 和 F(4, 3) 能大幅减少乘法次数代价是数值精度下降和额外变换开销。如果面试聊到性能优化能说出这一层是比单纯背 im2col 更高级的答案。我自己在实际项目中用到手撕 Conv2D 的场景更多是调试自定义算子或者验证量化误差。比如在写量化感知训练时需要核验某个自定义卷积实现是否真的和标准浮点结果一致这时候把这段朴素前向和反向代码当作“reference implementation”任务瞬间就清晰了。它不一定快但它绝对可靠是所有后续复杂操作的基础设施。如果你正在准备面试我建议你别只背代码把这篇里的映射关系、推导过程、对拍方法都亲手过一遍。纸上得来终觉浅这种需要脑内坐标系旋转的题目真得靠手写几遍才能形成肌肉记忆。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑