CUDA学习之路[12]:卷积计算详解

8371 字
42 分钟
CUDA学习之路[12]:卷积计算详解
更新中

写在前边#

完整的代码仓库在这里,希望对大家有所帮助和收获,希望大家多多Star!

卷积神经网络,作为 AI 时代的鼻祖,是我们学习深度学习绝对不能错过的一个点。很多小伙伴对卷积操作依然一知半解。关于卷积这一系列教程,我参考了北京邮电大学计算机学院孙其博教授的授课内容,在此仅作学习交流使用,如果涉及侵权,请联系我及时下架与删除。

卷积神经网络总览#

现代 AI 几乎就是从卷积神经网络(CNN)开始的。虽然现在大家都在谈 Attention、谈 Transformer,但严格意义上来说,CNN 才算是真正替人类开创了一个时代的功臣。

小故事课堂

在 AlexNet 横空出世之前,神经网络完全就是半死不活的状态。在 2000 年初,你要是敢在学术期刊或者会议论文上写 “XX Neural Network”,大概率会被审稿人秒拒。在当时,这基本跟现在宣称自己造出了“永动机”差不多。

当时学术界流行的是专家模型和各种手工设计的特征算子。而 AlexNet 的伟大之处,就在于它让大家第一次看到:直接用原始像素做输入,网络就可以自己提取出一整套惊艳的视觉特征,全程完全不需要人类手动设计。随后,目标检测、语义分割等一系列技术迅速跟上。不过风水轮流转,时至今日传统 CV 又沉寂下去了。只不过,这次不是因为被瞧不起,而是因为在当前框架下大家已经卷到头了,时代的红利暂时吃完了。

卷积,究竟解决了什么问题?#

回到正题,既然神经网络这么猛,那为什么以前的方案不行?在 AlexNet 出现之前,大部分网络采用的都是 MLP(多层感知机 / 全连接网络)去跑图像,这个方案有两个非常严重的致命伤:参数量爆炸 以及 缺乏几何直觉

而卷积操作之所以能拯救全连接网络的臃肿,是因为它比 Linear多了两个最关键的物理特性/先验假设:

  1. 局部相关性:图像中的一个像素,往往只和它周围的邻居有关系,和远处的像素无关。
  2. 平移不变性:一只猫无论出现在图像的左上角还是右下角,它都是一只猫。负责提取“猫耳朵”特征的算子,在图像的任意一个地方都可以通用。

正是因为有了这两个先验,卷积核(Kernel)应运而生。它在数学和工程上扮演的角色,其实就是滤波器(Filter)模板匹配器

卷积运算本质

卷积运算在行为上,是用一个卷积核在输入特征图上滑动,每个位置做逐元素乘积再求和,最终得到一个标量输出。

当卷积核滑过特征图时,如果某块区域的特征结构与卷积核的权重分布高度相似,卷积后的结果就会非常大(强激活);如果完全不匹配,结果就很小。这样,我们就能用一堆卷积核,在图片中大范围地搜索出若干特定的几何结构或视觉模式。

通信卷积和 CV 卷积的区别是什么?

实际上在通信领域(尤其是数字信号处理 DSP)中,卷积同样是基础概念。很多工科生看到深度学习里的卷积难免会产生一丝恍惚。

简单来说:通信中的卷积反映的是系统在时间上的“历史累积效应”(有因果律和信号衰减),而 CV 中的卷积则是空间上的“局部模式匹配”(纯粹的空间Filter)。

为什么网络必须是“多层”的?#

既然卷积可以通过滑动窗口完美匹配出局部特征,那问题来了:我们只用一个厚厚的卷积层,直接去识别整张图,到底行不行?

答案是:完全不行。这就引出了 CNN 核心设计中的另一个灵魂:多层级结构。因为人类认识这个世界的过程,本身就是具有高度层级性的:我们总是从 像素 \rightarrow 边的组合 \rightarrow 局部器官 \rightarrow 最终演进到完整的物体

现代的神经网络完美模拟了这一认知流:

  • 低层网络(靠近输入):感受野很小,只能看到几个像素。由于卷积核的参数是局部的,它只能提取极其基础的几何线条模式,例如亮暗边缘、45度角、基础色彩等。
  • 中层网络:负责将低层提取的“线条特征”进行拼装组合,这时候在网络眼里,它们变成了更加复杂的组合特征或局部部件(比如圆圈、网格、或者猫的耳朵、汽车的轮子)。
  • 高层网络(靠近输出):将中层的部件进一步组装,升级为具有全局语义信息的宏观概念(比如一张完整的猫脸,或者一辆整车)。
实际上有相关的论文专门探讨过深度和感受野哪一个对于网络更加有效,我记得是谷歌做的,大家可以查一下?

如果把这种“层级抽象”的直觉翻译成数学语言,你就能更直观地理解为什么“堆叠多层小核”比“单层大核”要高效得多。

在数学层面,单层网络的核心是:

y=σ(Wx+b)y = \sigma(Wx+b)

这只是一个简单的非线性变换,它能够在高维空间里切出来的“分类边界”是极其有限的,根本无法拟合现实世界复杂的物体流形。

多层网络的形式则是函数的连续复合:

y=σ(Wn...σ(W2σ(W1x+b1)+b2)...)y = \sigma(W_n ... \sigma (W_2 \sigma(W_1x+b_1)+b_2)...)

通过增加层数(深度),网络能够以惊人的参数效率,快速逼近极其复杂的、高度扭曲的高维非线性边界。

深度学习之所以能叫“深度”,核心就在于它用“空间的层数”换取了“认知的广度(感受野)”与“计算的效率(参数共享)”。

卷积神经网络#

最简单的卷积神经网络包含以下几个模块:

  • 卷积层
  • 激活层
  • 池化层
  • 全连接层
  • 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

上边就是一个最最简单的卷积神经网络。

其中核心是卷积层。

我们来看一下具体的卷积核的作用

两张示意图
两张示意图
对于这两张图,它们分别是RGB图与灰度图。

我可以定义一系列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+1
assert list(out.shape) == [1, 32, 30, 30]

我们看它的数学表达(以高度 OHOH 为例,宽度 OWOW 的计算完全同理):

OH=H+2pd(R1)1s+1OH = \lfloor \frac{H + 2p - d(R - 1) - 1}{s} \rfloor + 1

公式符号说明:

  • HH:输入特征图的高度
  • pp:单侧挂空白像素的层数(Padding)
  • dd:空洞率(Dilation),默认是 1
  • RR:卷积核的高度(Kernel Height)
  • ss:滑动步长(Stride)
  • \lfloor \dots \rfloor:向下取整符号(数学里的 Floor 函数,在 Python 里就是整除 //

首先输入特征图的宽和高、卷积核的宽和高都很好理解,但是padding、空洞率、滑动步长是什么?

padding的含义是我们在图像上下左右补一层空白,例如原本是3*3的特征图,如果padding=1就说明上下左右各补一层,就变成5*5了

啥是空洞卷积 (了解即可)

普通的卷积核它是实心的,例如3*3,它的实际大小就是3*3。 但是对于有空洞的,它卷积核内部元素之间会被塞入空心的像素,因此它实际在图像上覆盖的跨度变成了d(R1)+1d(R - 1) + 1

最后是滑动长度Stride,卷积核要在特征图上,从上到下从左到右依次滑动,而这个滑动的步长就代表Stride。

其实这个是更常见的计算方式,空洞卷积不常见

输出尺寸=输入尺寸+2×Padding有效核大小Stride+1\text{输出尺寸} = \frac{\text{输入尺寸} + 2 \times \text{Padding} - \text{有效核大小}}{\text{Stride}} + 1

在实际工业界搭建Backbone时,工程师们其实很少去天天算这种奇葩的不对称尺寸。为了让特征图宽高的变化符合预期,业界总结了几套固定的尺寸配方

尺寸减半(Downsampling)
  • 核心目的:代替池化层,让特征图的高宽缩水一半。
  • 通用组合Kernel=3, Stride=2, Padding=1
尺寸保持(Same Padding)
  • 核心目的:只做特征提取,不改变特征图的宽高分辨率。
  • 通用组合 1Kernel=3, Stride=1, Padding=1 \rightarrow 输出尺寸与输入完全一致。
  • 通用组合 2Kernel=5, Stride=1, Padding=2 \rightarrow 输出尺寸与输入完全一致。

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版本#

首先,我们要清楚inputweight的形状分别是:

  • Input: [B, C_in, H, W]
  • Weight: [C_out, C_in, kH, kW]
B, _, _, _ = input_tensor.shape
C_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 + 1
W_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 + 1
W_out = (input_padded.shape[3] - kW) // stride_w + 1
L = 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
  • 时间复杂度:O(B×Cout×Hout×Wout×Cin×Kh×Kw)O(B \times C_{out} \times H_{out} \times W_{out} \times C_{in} \times K_{h} \times K_{w})
  • 空间复杂度:O(1)O(1)

Naive原本不需要开辟额外的大块内存。它只加载当前窗口那一点数据,对缓存非常友好,但是可以看到它的时间复杂度爆掉了。

Im2col
  • 时间复杂度:理论上和上边一样,但是把卷积运算改为了大矩阵乘法,可以用GPU进行并行加速
  • 空间复杂度:O(B×Hout×Wout×Cin×Kh×Kw)O(B \times H_{out} \times W_{out} \times C_{in} \times K_h \times K_w)

im2col临时内存占用巨大,实际上相邻的滑动窗口是有极其大量重叠的像素的,im2col把这些像素强制复制并展开,所以很容易爆显存。

One more thing

思考一下如何优化im2col,提示:implicit GEMM convolution & Winograd。

当GEMM优化到一定程度,它就不再是一个Memory Bound的任务,转而变成了Compute Bound任务。因此后续的优化思路便是降低计算复杂度。

这几个算法本身比较复杂,大家可以自行扩展。

看一下两个版本的效果差异:

Naive & Official
Naive & Official

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中启动了 B×Cout×Hout×WoutB\times C_{out} \times H_{out} \times W_{out} 个程序实例,每一个负责计算一个输出像素。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端的判断条件,提前填充后所有的卷积窗口都是落在有效区域内的,内核并不需要额外的边界判断。

对于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 # 得到batch
4. 算出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版本有点类似。

  1. 计算感受野的起始位置
  2. 初始化累加器
  3. 在输入通道循环
  4. 卷积核空间循环
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?

实际上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
  1. 权重加载从每个输出点一次变成了每个Tile一次。
  2. 输入加载从标量变成了连续向量。
  3. 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 + kw
wgt_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=C_out?

从复用的角度来看,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的效果很明显,但代价是显式构建了一个巨大的中间矩阵,这个矩阵有多大呢?

输入张量:

x:[B,Cin,Hin,Win]x:[B, C_{in}, H_{in}, W_{in}]

F.unfold后得到的形状是:

col:[B,Cin×kH×kW,Hout×Wout]col: [B, C_{in} \times kH \times kW, H_{out} \times W_{out}]

经过permute和reshape,最终GEMM的形状为:

col2d:[Cin×kH×kW,B×Hout×Wout]col_2d: [C_{in} \times kH \times kW, B \times H_{out} \times W_{out}]

空间膨胀比为:

膨胀比=kHkWHoutWoutHinWin膨胀比 = \frac{kH \cdot kW \cdot H_{out} \cdot W_{out}}{H_{in} \cdot W_{in}}

如果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_out
b_idx = n_offs // hw # 批次索引
rem = n_offs % hw
h_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 // khw
rs = k_offs % khw
kh_ = rs // kW
kw_ = 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,而是通过数学运算降低了乘法,感兴趣的可以自行学习,由于笔者能力有限,在此不做赘述。

支持与分享

如果这篇文章对你有帮助,欢迎分享给更多人或赞助支持!

赞助
CUDA学习之路[12]:卷积计算详解
https://dlog.com.cn/posts/cuda12/conv/
作者
杜子源
发布于
2026-06-12
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author
杜子源
都是风景,幸会
公告
请狠狠地打赏我,打赏一次,爆更一篇!!
音乐
封面

音乐

暂未播放

0:00 0:00
暂无歌词
分类
标签
站点统计
文章
29
分类
8
标签
11
总字数
93,152
运行时长
0
最后活动
0 天前

目录