CUDA学习之路[12]:卷积计算详解
写在前边
完整的代码仓库在这里,希望对大家有所帮助和收获,希望大家多多Star!
卷积神经网络,作为 AI 时代的鼻祖,是我们学习深度学习绝对不能错过的一个点。很多小伙伴对卷积操作依然一知半解。关于卷积这一系列教程,我参考了北京邮电大学计算机学院孙其博教授的授课内容,在此仅作学习交流使用,如果涉及侵权,请联系我及时下架与删除。
卷积神经网络总览
现代 AI 几乎就是从卷积神经网络(CNN)开始的。虽然现在大家都在谈 Attention、谈 Transformer,但严格意义上来说,CNN 才算是真正替人类开创了一个时代的功臣。
在 AlexNet 横空出世之前,神经网络完全就是半死不活的状态。在 2000 年初,你要是敢在学术期刊或者会议论文上写 “XX Neural Network”,大概率会被审稿人秒拒。在当时,这基本跟现在宣称自己造出了“永动机”差不多。
当时学术界流行的是专家模型和各种手工设计的特征算子。而 AlexNet 的伟大之处,就在于它让大家第一次看到:直接用原始像素做输入,网络就可以自己提取出一整套惊艳的视觉特征,全程完全不需要人类手动设计。随后,目标检测、语义分割等一系列技术迅速跟上。不过风水轮流转,时至今日传统 CV 又沉寂下去了。只不过,这次不是因为被瞧不起,而是因为在当前框架下大家已经卷到头了,时代的红利暂时吃完了。
卷积,究竟解决了什么问题?
回到正题,既然神经网络这么猛,那为什么以前的方案不行?在 AlexNet 出现之前,大部分网络采用的都是 MLP(多层感知机 / 全连接网络)去跑图像,这个方案有两个非常严重的致命伤:参数量爆炸 以及 缺乏几何直觉。
而卷积操作之所以能拯救全连接网络的臃肿,是因为它比 Linear多了两个最关键的物理特性/先验假设:
- 局部相关性:图像中的一个像素,往往只和它周围的邻居有关系,和远处的像素无关。
- 平移不变性:一只猫无论出现在图像的左上角还是右下角,它都是一只猫。负责提取“猫耳朵”特征的算子,在图像的任意一个地方都可以通用。
正是因为有了这两个先验,卷积核(Kernel)应运而生。它在数学和工程上扮演的角色,其实就是滤波器(Filter)或模板匹配器。
卷积运算在行为上,是用一个卷积核在输入特征图上滑动,每个位置做逐元素乘积再求和,最终得到一个标量输出。
当卷积核滑过特征图时,如果某块区域的特征结构与卷积核的权重分布高度相似,卷积后的结果就会非常大(强激活);如果完全不匹配,结果就很小。这样,我们就能用一堆卷积核,在图片中大范围地搜索出若干特定的几何结构或视觉模式。
实际上在通信领域(尤其是数字信号处理 DSP)中,卷积同样是基础概念。很多工科生看到深度学习里的卷积难免会产生一丝恍惚。
简单来说:通信中的卷积反映的是系统在时间上的“历史累积效应”(有因果律和信号衰减),而 CV 中的卷积则是空间上的“局部模式匹配”(纯粹的空间Filter)。
为什么网络必须是“多层”的?
答案是:完全不行。这就引出了 CNN 核心设计中的另一个灵魂:多层级结构。因为人类认识这个世界的过程,本身就是具有高度层级性的:我们总是从 像素 边的组合 局部器官 最终演进到完整的物体。
现代的神经网络完美模拟了这一认知流:
- 低层网络(靠近输入):感受野很小,只能看到几个像素。由于卷积核的参数是局部的,它只能提取极其基础的几何线条模式,例如亮暗边缘、45度角、基础色彩等。
- 中层网络:负责将低层提取的“线条特征”进行拼装组合,这时候在网络眼里,它们变成了更加复杂的组合特征或局部部件(比如圆圈、网格、或者猫的耳朵、汽车的轮子)。
- 高层网络(靠近输出):将中层的部件进一步组装,升级为具有全局语义信息的宏观概念(比如一张完整的猫脸,或者一辆整车)。
如果把这种“层级抽象”的直觉翻译成数学语言,你就能更直观地理解为什么“堆叠多层小核”比“单层大核”要高效得多。
在数学层面,单层网络的核心是:
这只是一个简单的非线性变换,它能够在高维空间里切出来的“分类边界”是极其有限的,根本无法拟合现实世界复杂的物体流形。
而多层网络的形式则是函数的连续复合:
通过增加层数(深度),网络能够以惊人的参数效率,快速逼近极其复杂的、高度扭曲的高维非线性边界。
深度学习之所以能叫“深度”,核心就在于它用“空间的层数”换取了“认知的广度(感受野)”与“计算的效率(参数共享)”。
卷积神经网络
最简单的卷积神经网络包含以下几个模块:
- 卷积层
- 激活层
- 池化层
- 全连接层
- Softmax输出头
- …
class Net(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 10, 5) self.pool = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(10, 20, 5) self.fc = nn.Linear(320, 10)
def forward(self, x): x = self.pool(torch.relu(self.conv1(x))) x = self.pool(torch.relu(self.conv2(x))) x = x.view(-1, 320) x = self.fc(x) return x上边就是一个最最简单的卷积神经网络。
其中核心是卷积层。
我们来看一下具体的卷积核的作用

我可以定义一系列3*3的卷积核:

具体卷积核的样式如下:
kernels = { 'Sobel X': torch.tensor([[-1., 0., 1.], [-2., 0., 2.], [-1., 0., 1.]]), 'Sobel Y': torch.tensor([[-1.,-2.,-1.], [ 0., 0., 0.], [ 1., 2., 1.]]), 'Laplacian': torch.tensor([[ 0., 1., 0.], [ 1.,-4., 1.], [ 0., 1., 0.]]), 'Gaussian': torch.tensor([[1.,2.,1.], [2.,4.,2.], [1.,2.,1.]]) / 16., 'Sharpen': torch.tensor([[ 0.,-1., 0.], [-1., 5.,-1.], [ 0.,-1., 0.]]), 'Emboss': torch.tensor([[-2.,-1., 0.], [-1., 1., 1.], [ 0., 1., 2.]]),}它们的功能和它们的名字一样,例如有的是提取水平边缘,有的是提取垂直边缘,有的则是进行高斯模糊,我们来具体看一下这些卷积核作用在特征图上的效果:

以Sobel X为例,很明显,它的垂直的头发那一侧的线条更加明显,而Sobel Y中水平的线条则更加明显。 Gaussian则是可以看出相比较原图变得更加模糊了。
这些卷积核都有对应的一些特征。而这些特征是我们人为定义好的特征核。 但是有了AI之后,或者有了前向传播和反向传播之后,这些卷积核都可以自动来进行训练。
就是对若干卷积核、权重不断进行调整,使其满足精度要求。
Conv2d API与张量布局
PyTorch 提供两套等价接口:
| 接口 | 形式 | 适用场景 |
|---|---|---|
torch.nn.functional.conv2d | 函数式,权重显式传入 | 临时计算、自定义层、教学验证 |
torch.nn.Conv2d | 模块式,权重作为可学习参数 | 搭建 nn.Module 网络 |
二者底层调用同一套实现(cuDNN / MKLDNN)。
张量布局 (PyTorch 默认 NCHW):
input : [N, C, H, W] batch × 输入通道 × 高 × 宽weight: [K, C, R, S] 输出通道 × 输入通道 × 核高 × 核宽bias : [K] 每输出通道一个偏置(可选)output: [N, K, OH, OW] OH = floor((H + 2p - d(R-1) - 1)/s) + 1我们先看一下最基本的调用:
inp = torch.randn(1, 16, 32, 32) # [N=1, C=16, H=32, W=32]wt = torch.randn(32, 16, 3, 3) # [K=32, C=16, R=3, S=3]
# weight 的 C 必须与 input 的 C 一致out = F.conv2d(inp, wt, stride=1, padding=0)
print(f"input : {list(inp.shape)}") # [1, 16, 32, 32]print(f"weight: {list(wt.shape)}") # [32, 16, 3, 3]print(f"output: {list(out.shape)}") # [1, 32, 30, 30] 30 = 32-3+1assert list(out.shape) == [1, 32, 30, 30]我们看它的数学表达(以高度 为例,宽度 的计算完全同理):
公式符号说明:
- :输入特征图的高度
- :单侧挂空白像素的层数(Padding)
- :空洞率(Dilation),默认是 1
- :卷积核的高度(Kernel Height)
- :滑动步长(Stride)
- :向下取整符号(数学里的 Floor 函数,在 Python 里就是整除
//)
首先输入特征图的宽和高、卷积核的宽和高都很好理解,但是padding、空洞率、滑动步长是什么?
padding的含义是我们在图像上下左右补一层空白,例如原本是3*3的特征图,如果padding=1就说明上下左右各补一层,就变成5*5了
普通的卷积核它是实心的,例如3*3,它的实际大小就是3*3。 但是对于有空洞的,它卷积核内部元素之间会被塞入空心的像素,因此它实际在图像上覆盖的跨度变成了。
最后是滑动长度Stride,卷积核要在特征图上,从上到下从左到右依次滑动,而这个滑动的步长就代表Stride。
在实际工业界搭建Backbone时,工程师们其实很少去天天算这种奇葩的不对称尺寸。为了让特征图宽高的变化符合预期,业界总结了几套固定的尺寸配方。
- 核心目的:代替池化层,让特征图的高宽缩水一半。
- 通用组合:
Kernel=3, Stride=2, Padding=1
- 核心目的:只做特征提取,不改变特征图的宽高分辨率。
- 通用组合 1:
Kernel=3, Stride=1, Padding=1输出尺寸与输入完全一致。 - 通用组合 2:
Kernel=5, Stride=1, Padding=2输出尺寸与输入完全一致。
Convxd的Pytorch接口
通常来说,常见的卷积操作有3中:Conv1d、Conv2d、Conv3d,它们的使用方式如下:
conv1d = nn.Conv1d(in_channel, out_channels=3, kernel_size=3)conv2d = nn.Conv2d(in_channel, out_channels=3, kernel_size=(H, W))conv3d = nn.Conv3d(in_channel, out_channels=3, kernel_size=(D, H, W))卷积核在一维中,通常主要处理时间数据,因此只需要向左或者向右移动;在二维中,内核会向两个方向移动,因此除了左右一般还有上下;而在三维卷积核中,内核就可以在三维空间中滑动。一般来说三维卷积常在医学领域中用到(X光片的切片等等)
这里指的并不是处理的输入数据,而是卷积核在数据上的移动轨迹。一维卷积核也可以在二维图像上处理数据。
后续我们的实现也主要以二维为主。
Pytorch版本手动实现
在Pytorch中,我主要实现naive版本和im2col版本,其中naive版本在上边有视频做演示,我们直接来看是如何实现的。
Naive版本
首先,我们要清楚input和weight的形状分别是:
- Input:
[B, C_in, H, W] - Weight:
[C_out, C_in, kH, kW]
B, _, _, _ = input_tensor.shapeC_out, _, kH, kW = weight.shape
stride_h = stride if isinstance(stride, int) else stride[0]stride_w = stride if isinstance(stride, int) else stride[1]
pad_h = padding if isinstance(padding, int) else padding[0]pad_w = padding if isinstance(padding, int) else padding[1]
if pad_h > 0 or pad_w > 0: input_padded = F.pad(input_tensor, (pad_w, pad_w, pad_h, pad_h), mode='constant', value=0.0)else: input_padded = input_tensor
_, _, H_pad, W_pad = input_padded.shape上边做的主要工作是把stride、pad都从参数中提取出来,然后把原本的input进行四周的填充。 填充好之后计算输出的宽和高:
# 计算输出的宽高尺寸H_out = (H_pad - kH) // stride_h + 1W_out = (W_pad - kW) // stride_w + 1计算完成之后就开始根据输出进行遍历:
for b in range(B): # 遍历 Batch 中的每一张图 for co in range(C_out): # 遍历每一个卷积核 for h in range(H_out): # 遍历输出特征图的高 for w in range(W_out): # 遍历输出特征图的宽
# 上述for循环都是从输出的视角计算的 # 输出特征图的每一个像素,都对应输入图像上的一个窗口
h_start = h * stride_h # 窗口在高度方向的起始行 w_start = w * stride_w # 窗口在宽度方向的起始行 h_end = h_start + kH # 窗口在高度方向的结束行(不包含) w_end = w_start + kW # 窗口在宽度方向的结束行(不包含)
# 提取当前窗口内所有输入通道的数据,形状为 (C_in, kH, kW) current_window = input_padded[b, :, h_start:h_end, w_start:w_end]
# 获取当前输出通道对应的卷积核权重,形状为 (C_in, kH, kW) current_kernel = weight[co]
# pixel_value = torch.sum(current_window * current_kernel) pixel_value = torch.einsum('cij,cij->', current_window, current_kernel)
# 如果有偏置,加上偏置 if bias is not None: pixel_value += bias[co]
# 将计算的标量写入到对应位置 output[b, co, h, w] = pixel_value其中从外到内依次是,B表示图片个数,co表示通道数,h, w表示的则是宽和高对应的输出,可以看到这里有4个for循环,执行起来必然是很慢的。
其中内部的核心就是:输出是一个标量,是根据一个滑动窗口决定的,因此要先计算出这个窗口的四个角的坐标。
计算好窗口的位置之后,直接与卷积核进行逐元素乘加即可。(看到这里可以想一下Triton和CUDA的优化思路)
最终把计算好的标量写入到对应的位置即可。
im2col版本
im2col 就是 image to column 的缩写,它是一种将输入特征图转换为矩阵的算法。im2col 算法的主要思想是将卷积核在输入特征图上滑动,每次滑动的步长为卷积核的步长,然后将卷积核覆盖的区域拉成一个列向量,最后将所有的列向量拼接在一起,就得到了一个矩阵。
下面我们通过一个简单的例子来说明 im2col 算法的原理:
# im2col的核心差异# [B, C_in * kH * kW, H_out * W_out]cols = F.unfold(input_padded, kernel_size=(kH, kW), stride=(stride_h, stride_w))
H_out = (input_padded.shape[2] - kH) // stride_h + 1W_out = (input_padded.shape[3] - kW) // stride_w + 1L = H_out * W_out
W_mat = weight.view(C_out, -1) # [C_out, C_in * kH * kW]
# [B, L, C_in*kH*kW] * [C_in*kH*kW, C_out]out = torch.matmul(cols.transpose(1, 2), W_mat.transpose(0, 1))# 输出形状:[B, L, C_out]out = out.transpose(1, 2)
if bias is not None: out += bias.view(1, -1, 1)
out = out.view(B, C_out, H_out, W_out)相比较Naive版本来说,im2col把原本的多个for循环压缩成为了一次大的矩阵乘法计算
两个版本的性能对比分析
- 时间复杂度:
- 空间复杂度:。
Naive原本不需要开辟额外的大块内存。它只加载当前窗口那一点数据,对缓存非常友好,但是可以看到它的时间复杂度爆掉了。
- 时间复杂度:理论上和上边一样,但是把卷积运算改为了大矩阵乘法,可以用GPU进行并行加速
- 空间复杂度:
im2col临时内存占用巨大,实际上相邻的滑动窗口是有极其大量重叠的像素的,im2col把这些像素强制复制并展开,所以很容易爆显存。
思考一下如何优化im2col,提示:implicit GEMM convolution & Winograd。
当GEMM优化到一定程度,它就不再是一个Memory Bound的任务,转而变成了Compute Bound任务。因此后续的优化思路便是降低计算复杂度。
这几个算法本身比较复杂,大家可以自行扩展。
看一下两个版本的效果差异:

Triton版本手动实现
Naive版本实现
前边我们描述卷积时,写出了这样一段代码:
for b in range(B): for co in range(C_out): for h in range(H_out): for w in range(W_out): # 取窗口、取卷积核、乘加、写入一个像素这个是由外到内,按照顺序逐像素进行计算。
那么我们要思考一个问题,这些循环之间,例如B和B+1,co和co+1,h和h+1,w和w+1这些迭代之间,有没有依赖关系?
答案:没有。
输出特征图的每一个像素output[b, co, h, w]的计算只依赖输入图像上的一个局部窗口(由b,h,w决定)以及某一个卷积核(由co决定)。
它并不依赖其他输出像素是否已经算完,也不依赖其他输出像素值,每个输出像素的计算都是完全独立的。
这就是并行化的基础。
在我们学习到了卷积之后,下意识的会按照“卷积核从左到右、从上到下滑动”这种动态视角去看问题,但是我们需要转变视角,而是站在每一个输出像素去思考。 把原本滑动定义为:一个输出点对应一个独立任务。
有了这种并行化思路,那他可以归结到我们之前学过的:element wise。
接下来就是如何实现并行化,以及如何优化的问题了。
原本是四层for循环,这里我们可以把它展开:
grid = (B * C_out * H_out * W_out,)这里表明,我们在Triton中启动了 个程序实例,每一个负责计算一个输出像素。grid网格大小正好等于输出元素的总数。
每个线程拿到自己的pid,反推出自己负责的(b, co, h, w),然后完成该输出点的卷积累加。
具体来说,我们先把封装函数弄明白:
def conv_v0(x, w, bias=None, stride=1, padding=0): # 1. 获取输入、权重的形状信息 B, C_in, H_in, W_in = x.shape C_out, _, kH, kW = w.shape
# 2. 解析步长与填充参数 sh, sw, ph, pw = parse_conv_params(stride, padding)
# 3. 对输入进行填充 (pad) x_pad = pad_input(x, ph, pw) _, _, H_pad, W_pad = x_pad.shape
# 4. 计算输出的空间尺寸 H_out, W_out = compute_output_dims(H_in, W_in, kH, kW, sh, sw, ph, pw)
# 5. 预分配输出张量 out = torch.empty(B, C_out, H_out, W_out, device=x.device, dtype=torch.float32)
# 6. 设定grid 并启动内核 grid = (B * C_out * H_out * W_out,) _kernel[grid]( x_pad, w, out, bias if bias is not None else None, B, C_in, H_pad, W_pad, C_out, kH, kW, sh, sw, H_out, W_out ) return out对于封装函数来说,主要是对输入填充,以及提前分配好存储,最后再启动内核进行计算。
简化kernel端的判断条件,提前填充后所有的卷积窗口都是落在有效区域内的,内核并不需要额外的边界判断。
对于Kernel代码,首先进行的是传参或者接口:
def _kernel(input_ptr, weight_ptr, output_ptr, bias_ptr, B, C_in, H_in, W_in, C_out, kH, kW, sh, sw, H_out, W_out):我之前已经提到过很多次了,要把input和output指针传入,但是除了这些之外,还需要传入和stride相关的,这是为什么呢?——为了防止传入到kernel中的低层内存不连续。
我们进行传参之后,先要将线程id映射到输出坐标。在映射之前,有一个默认规定,那就是张量在内存中是行主序,维度顺序由内向外依次为:
最内层 -》 宽度 -》 高度 -》 通道 -》 Batch也就是说,如果我们把一个4维坐标(b, co, h, w)展成一维索引,公式就是:
pid = b * (C_out * H_out * W_out) + co * (H_out * W_out) + h * W_out + w接下来就是逆向推导
已知3721秒,请问这是X时X分X秒。
已知:pid是一个已经算好的数待求:(b, co, h, w)
1. 算出宽度坐标 w = pid % W_out rem = pid // W_out # 得到了剩下的(高度、通道、batch)2. 算出高度坐标 h = rem % H_out rem = rem // H_out # 得到了(通道、batch)3. 算出通道 co = rem % C_out rem = rem // C_ot # 得到batch4. 算出batch b = rem pid = tl.program_id(0) # 每个block都有一个pid w = pid % W_out # 宽度坐标 rem = pid // W_out # 剩下的部分 h = rem % H_out # 高度坐标 rem = rem // H_out # 剩下的部分 co = rem % C_out # 输出通道坐标 b = rem // C_out # batch坐标既然已经算出来了输出坐标,那就要对应计算输入窗口的起始位置。
这个地方就和原本的Pytorch的naive版本有点类似。
- 计算感受野的起始位置
- 初始化累加器
- 在输入通道循环
- 卷积核空间循环
h_in = h * sh; w_in = w * sw acc = tl.zeros([], dtype=tl.float32)
for ci in range(C_in): in_c_base = b * C_in * H_in * W_in + ci * H_in * W_in wgt_c_base = co * C_in * kH * kW + ci * kH * kW for kh in range(kH): for kw in range(kW): in_idx = in_c_base + (h_in + kh) * W_in + (w_in + kw) wgt_idx = wgt_c_base + kh * kW + kw acc += tl.load(input_ptr + in_idx) * tl.load(weight_ptr + wgt_idx)
if bias_ptr is not None: acc += tl.load(bias_ptr + co)
out_idx = b * C_out * H_out * W_out + co * H_out * W_out + h * W_out + w tl.store(output_ptr + out_idx, acc)这个版本的并行,实际上是针对每个输出点进行并行,将每个输出点映射到GPU线程中,原本的(b, co, h, w)被GPU的线程grid一次性替代了,实现了从O(N)到O(1)的加速,其中N是输出元素的总数。
实际上offset和mask是为了元素定位,现在我们通过手动计算(b, co, h, w)已经进行了实现,并且输出尺寸完全匹配grid,所以并不需要mask进行屏蔽写回。
Spatial版本实现
Triton已经解决了最基础的并行:每个线程分别计算所有的输出点。
一个线程只负责一个输出点,却要遍历C_in \times kH \times kW次全局访问,计算密度极低,访存明显成为瓶颈。
- 对于权重来说,相邻输出点用到的kernel是完全相同的,例如(h, w)和(h, w+1)用的都是同一组权重,但是在v0中它们属于两个完全独立的线程,各自把所有的权重重新从全局内存加载一遍;
- 对于输入来说,输入重复读取。以
3\times 3核为例,输入图一个像素被9个不同的输出点用到,但是在v0中,这9个输出点分别属于9个线程,彼此不共享数据,导致同一个输入像素被反复加载9次。
v0版本的加载都是标量加载,一次只load一个float,线程访存地址是极其离散的,无法利用GPU的合并访存机制。
因此,v1的思路主要是两个:数据复用+向量加载。
v1版本不再让一个线程只计算一个输出点,而是让它负责输出图上的二维特征块,
实际上网格是3D的(n_tiles_h * n_tiles_w, C_out, B), 一个线程块负责输出通道co,批次b里的一个空间块,大小为TILE_H \times TILE_W。
也就是说,一次性计算256个输出像素,而不是计算单个像素。
在kernel内部,原本一维的加载现在改成二维:
h_offs = tile_h * TILE_H + tl.arange(0, TILE_H) # 形状 (TILE_H,) w_offs = tile_w * TILE_W + tl.arange(0, TILE_W) # 形状 (TILE_W,) acc = tl.zeros([TILE_H, TILE_W], dtype=tl.float32)然后对每一个ci, kh, kw,我们也不再循环标量,而是进行向量化块加载:
in_h = h_offs * sh + kh # (TILE_H,) in_w = w_offs * sw + kw # (TILE_W,)
# 一次性加载一整个输入块 in_vals = tl.load( input_ptr + ... + in_h[:, None] * W_in + in_w[None, :], mask=..., other=0.0 ) # 形状 (TILE_H, TILE_W)
# 只加载一个权重标量 wgt_val = tl.load(weight_ptr + wgt_idx) # 标量
# 整个输出块的一次累加 acc += in_vals * wgt_val- 权重加载从每个输出点一次变成了每个Tile一次。
- 输入加载从标量变成了连续向量。
- tile内部输入数据共享。
性能对比如下:
Kernel Time (ms) GFLOPS Peak Mem Status ──────────────────────────────────────────────────────────────────────── v0-naive 147.2876 23.4 12.0 ✓ v1-空间分块 0.9017 3815.1 12.0 ✓Channel版本实现
数据复用+向量加载是v1版本的精髓,也是所有算子优化的精髓,但是实则还有一块并没有处理:不同输出通道的计算是彼此独立的,但是它们用到的输入是完全相同的。
什么意思呢? 我们来回顾一下v1的grid:
grid = (n_tiles_h * n_tiles_w, C_out, B)每一个线程块固定处理一个输出通道co。这意味着对于同一个输入像素块,它被不同co的线程块各自独立地从全局内存加载了一遍。
例如:
输入:(1, 64, 56, 56)卷积核:3*3输出通道:64步长:1
输出:(1, 64, 54, 54)在 v1 中,要计算位置 (tile_h=0, tile_w=0) 的输出块,需要启动 64 个线程块。这 64 个块在循环到 ci=0, kh=0, kw=0 时,加载的是完全相同的输入像素块,都是从 ci=0 的输入图上抠出来的同一块 TILE_H × TILE_W 区域。 换句话说,同一份输入数据被 64 个 block 各读一次,全局内存中该输入块的访问量放大了 64 倍。
实际上输入只需要读取一次,同时给多个输出通道:
grid = (n_tiles_h * n_tiles_w, ceil(C_out / BLOCK_K), B)第二个维度并不是单个的C_out,而是切成BLOCK_K大小的通道块,一个线程块同时负责多个输出通道。
累加器从2D升级为3D:
k_offs = pid_k * BLOCK_K + tl.arange(0, BLOCK_K) # 通道偏移向量 (BLOCK_K,)acc = tl.zeros([TILE_H, TILE_W, BLOCK_K], dtype=tl.float32) # 三维累加器输入只需要加载一次:
in_2d = tl.load( input_ptr + in_ci_base + in_h[:, None] * W_in + in_w[None, :], mask=in_load_mask, other=0.0) # 形状仍是 [TILE_H, TILE_W]权重一次性加载BLOCK_K个通道的权重:
wgt_idx = k_offs * wgt_stride_k + wgt_ci_base + kh * kW + kwwgt_1d = tl.load(weight_ptr + wgt_idx, mask=k_mask, other=0.0) # [BLOCK_K]最终使用广播指令一次性完成所有通道的乘加:
acc += in_2d[:, :, None] * wgt_1d[None, None, :]# in_2d 扩展为 [TILE_H, TILE_W, 1]# wgt_1d 扩展为 [1, 1, BLOCK_K]# 结果自动广播为 [TILE_H, TILE_W, BLOCK_K]从复用的角度来看,BLOCK_K越大越好,但是线程块的资源是有限的,会导致寄存器占用量飙升,如果BLOCK_K太大,会导致隐藏延迟的能力反而下降。因此最好是进行autotune,探索最优组合并且限制它们的大小。
性能对比如下:
Kernel Time (ms) GFLOPS Peak Mem Status ──────────────────────────────────────────────────────────────────────── v0-naive 147.2876 23.4 12.0 ✓ v1-空间分块 0.9017 3815.1 12.0 ✓ v2-空间+通道分块 0.6692 5139.9 12.0 ✓im2col版本实现
同理,在Pytorch版本我就说过,卷积本质上可以写成矩阵乘法。
def conv_v3(x, w, bias=None, stride=1, padding=0): B, C_in, H_in, W_in = x.shape C_out, _, kH, kW = w.shape sh, sw, ph, pw = parse_conv_params(stride, padding) H_out, W_out = compute_output_dims(H_in, W_in, kH, kW, sh, sw, ph, pw) CRS = C_in * kH * kW # 每个滑动窗口的特征长度核心就是算出输出空间尺寸,以及卷积核展开后的列维度,然后把输入图变成展开矩阵:
col = F.unfold(x, kernel_size=(kH, kW), padding=(ph, pw), stride=(sh, sw))这是 PyTorch 内置的 im2col 操作。F.unfold 会遍历输入图的每一个滑动窗口(包括 padding),把每个窗口内的像素按通道优先顺序拉直成一个列向量。
于是 col 是一个形状 [B, CRS, H_out*W_out] 的张量。每个列就对应一个输出位置所需的全部输入像素(共 CRS 个标量),按 (ci, kh, kw) 顺序排好。
之后进行矩阵运算,然后再变回去即可。
这里代码十分简洁,参考代码库即可。
im2col的效果很明显,但代价是显式构建了一个巨大的中间矩阵,这个矩阵有多大呢?
输入张量:
F.unfold后得到的形状是:
经过permute和reshape,最终GEMM的形状为:
空间膨胀比为:
如果stride和padding能够使得空间尺寸恰好不变,那么膨胀比就是kH · kW,一个3*3的卷积会让显存膨胀9倍,7*7的卷积会膨胀49倍。
性能对比和显存占用如下:
Kernel Time (ms) GFLOPS Peak Mem Status ──────────────────────────────────────────────────────────────────────── v0-naive 147.2876 23.4 12.0 ✓ v1-空间分块 0.9017 3815.1 12.0 ✓ v2-空间+通道分块 0.6692 5139.9 12.0 ✓ v3-im2col+GEMM 0.5541 6208.4 228.4 ✓Implicit GEMM版本实现
既然我们要用到一个巨大的中间临时矩阵,那么有没有什么办法取消这个显式展开呢?最好是能够直接从原始输入张量中读取。
这就提到了一个思想:指针,或者叫做视图。
如果我们有一种映射方案,能够把原图和中间图之间,通过函数映射起来,这样原本的开销就可以完全省下来。
v4 的 kernel 本质上在模拟矩阵乘法 [C_out, CRS] × [CRS, B*H_out*W_out],但它 不提前铺平 [CRS, B*H_out*W_out] 矩阵,而是:
- 在需要加载一个输入元素时,根据输出位置 n 和归约维度 k 反算出它在原始张量中的坐标,
- 直接从原始输入张量 input_ptr 中读取。
这样,整个卷积过程占用的内存只有:输入 x、权重 w、输出 out。
首先对于grid:
grid = (triton.cdiv(C_out, BLOCK_M), triton.cdiv(N_total, BLOCK_N))Kernel是2D网络,第0维沿着输出通道切块,第1维沿着输入位置切块,正好对应矩阵乘法的分块方案。
对于每个block,需要知道它负责的输出位置以及输入图对应的坐标:
n_offs = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) # [BLOCK_N]hw = H_out * W_outb_idx = n_offs // hw # 批次索引rem = n_offs % hwh_in = (rem // W_out) * sh # 输出位置对应的输入起始 h 坐标w_in = (rem % W_out) * sw # 输出位置对应的输入起始 w 坐标矩阵乘法的内层是沿着k=CRS进行规约,把v4切成BLOCK_K大小的块:
for k_start in range(0, CRS, BLOCK_K): k_offs = k_start + tl.arange(0, BLOCK_K) # [BLOCK_K] k_mask = k_offs < CRS每个k对应一个(ci, kh, kw)三元组,同样通过整除和取模恢复:
ci = k_offs // khwrs = k_offs % khwkh_ = rs // kWkw_ = rs % kW权重一次性加载BLOCK_M * BLOCK_K的块:
wgt = tl.load(weight_ptr + m_offs[:, None] * CRS + k_offs[None, :], mask=m_mask[:, None] & k_mask[None, :], other=0.0)对于输入来说,则是动态计算出来的:
in_h = h_in[None, :] + kh_[:, None] # [BLOCK_K, BLOCK_N]in_w = w_in[None, :] + kw_[:, None] # [BLOCK_K, BLOCK_N]in_idx = (b_idx[None, :] * C_in * H_in * W_in + ci[:, None] * H_in * W_in + in_h * W_in + in_w)inp = tl.load(input_ptr + in_idx, mask=k_mask[:, None] & n_mask[None, :], other=0.0)虽然有计算索引开销,但是避免了显式的数据重排带来的带宽浪费。
完整的性能对比如下:
Kernel Time (ms) GFLOPS Peak Mem Status ──────────────────────────────────────────────────────────────────────── v0-naive 147.2876 23.4 12.0 ✓ v1-空间分块 0.9017 3815.1 12.0 ✓ v2-空间+通道分块 0.6692 5139.9 12.0 ✓ v3-im2col+GEMM 0.5541 6208.4 228.4 ✓ v4-ImplicitGEMM 0.0850 40453.2 24.0 ✓ torch(cuDNN) 0.0988 34810.7 36.0 —CUDA版本
这里并没有实现CUDA版本,一来是和矩阵乘法有些类似,二来是Conv实际上有更加高效的算法,既不是im2col,也不是naive,而是通过数学运算降低了乘法,感兴趣的可以自行学习,由于笔者能力有限,在此不做赘述。
支持与分享
如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!