尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

GCNN解析:群等变卷积如何让CNN旋转鲁棒

GCNN解析:群等变卷积如何让CNN旋转鲁棒 我最近重新翻了一遍 Taco Cohen 和 Max Welling 的经典论文Group Equivariant Convolutional Networks这篇 ICML 2016 的工作某种程度上改变了我对 CNN 设计方式的理解。它做的事情用一句话概括把卷积神经网络从“只对平移等变”扩展到了“对旋转等变”让 CNN 不再只是“在图片里找边缘找纹理”而是“无论目标转了多少角度我都能用同一套参数把它认出来”。如果你正在做图像分类、目标检测或者对“如何把先验结构塞进网络”这件事感兴趣这篇论文非常值得精读。这篇笔记我想换个方式写不按论文原文顺序复述而是按“我读它的时候是怎么一步步想通的”来展开先从 CNN 遇到的旋转痛点说起然后过一遍群卷积的数学直觉再到网络结构和实验表现最后给出我复现时的代码思路和踩过的坑。1. 论文背景与核心问题1.1 CNN 对旋转变换为什么这么弱标准 CNN 最成功的设计之一就是局部卷积核加平移共享。一个竖直方向的边缘检测器不管它出现在图片左上角还是右下角同一个卷积核都能把它检测出来这种性质叫“平移等变性”。可以说CNN 天然内置了“目标位置不影响识别”的归纳偏置这也是它在图像任务上远好于全连接网络的重要原因。但问题也来了卷积核只能在一个方向、一个尺度上工作。你训练了一个检测竖直边缘的核当图像里的边缘旋转 90 度变成水平方向时这个核的输出会迅速衰减。换句话说CNN 对旋转变换是没有等变性的输入旋转了特征图不会像旋转后的“同样结构”那样继续被响应。有人可能会说数据增强不就行了吗在训练集里把每张图旋转几个角度再喂给网络网络不就能适应旋转了吗这个思路确实有效但它最多算“被动防御”。数据增强让网络在不同角度下分别学一遍特征本质上是在同一个任务上重复劳动网络并没有真正获得“旋转后结构不变”的抽象能力。后果就是参数利用率低训练时间变长而且在训练时没见过的旋转角度比如 17 度这种非规则角度上模型的泛化能力仍然有限。GCNN 的思路则完全不同它想在网络结构里就做出保证让特征提取对旋转变换是结构性的等变而不是靠堆数据硬怼。1.2 等变性和不变性的区别以及我们到底想要什么如果你已经接触过相关概念应该知道“等变”和“不变”是两码事。以人脸识别为例我们希望一个分类器无论输入的是人脸正脸还是旋转 30 度后的侧脸输出的身份标签都一样这叫“不变性”。但在网络的中间层我们其实不希望所有特征都是不变的——如果中间特征对旋转完全不变那就意味着网络丢失了“方向”这个信息。可是检测一个倾斜的鼻子、一个横着的嘴巴真的不需要区分方向吗显然需要。更合理的做法是中间层保持等变性输入人脸旋转了多少度中间特征图就以某种可预测的方式跟着“转”过去只有在最后一层再显式地把所有旋转方向的特征合并起来丢弃方向信息输出不变分类结果。这就像你让人去认一把椅子他围着椅子转一圈看到的每个侧面都不一样但这些不同的视角只是“椅子转了角度”不是“变成了一把新椅子”。等变性要的就是这种“视角变化可追踪”的性质。GCNN 的核心贡献就是把这种“可追踪”从平移扩展到了离散旋转群并且用群卷积统一实现了它。2. 群论基础与 GCNN 的核心思想2.1 从平移群到旋转群群到底是什么群论听起来吓人但你要是已经知道卷积是怎么在二维平面上“平移”的那群论离你只有一步之遥。一个群就是一个集合带上一种运算满足四条性质封闭性、结合律、存在单位元、每个元素都有逆元。平移操作本身就能构成群先平移 3 个像素再平移 5 个像素等于平移 8 个像素平移 0 个像素是单位元平移了 10 个像素平移 -10 就是它的逆。而且这些操作是封闭的组合起来仍然是平移操作。这篇论文里最重要的一类群是 P4 群。它包含两种操作一是平移二是围绕整数格点做 90 度整数倍的旋转。P4 群里的每个元素都可以写成“先旋转 r 次再平移 (u, v)”的复合操作。这类“平移和旋转组合起来”的群在数学上叫半直积用符号表示就是 P4 Z² ⋊ C4其中 C4 是 90 度旋转生成的循环群。如果你再把镜像反射也加进去就得到 P4M 群它额外包含水平翻转或垂直翻转这类操作。为什么要限定 90 度的整数倍而不是连续旋转最直接的原因是计算可行性离散群上的卷积可以严格表示为有限求和而连续旋转群需要复杂的测度理论和不可约表示2016 年那会儿还很难在深度学习框架里高效实现。而 90 度旋转已经能覆盖大量自然图像的对称模式比如方形网格、建筑物、家具。P4M 群则更进一步把镜像对称也拉进来适合那种“目标本身没有手性”的场景。2.2 群卷积的定义把“平移核”变成“变换核”标准 CNN 的卷积是这样写的(f * ψ)(t) Σ_{x} f(x) ψ(x - t)这里 t 是平移量ψ 是卷积核。你能看到卷积的本质是把核平移 t然后跟输入做内积。这个定义里唯一被允许的“核变换”就是平移。GCNN 的推广非常自然既然我们要对旋转等变那就应该允许核被旋转。于是群卷积写成这样(f * ψ)(g) Σ_{h∈G} f(h) ψ(g^{-1} h)注意这里的 f 和 ψ 都定义在群 G 上自变量从“平面坐标”变成了“群元素”。对于 P4 群群元素可以表示为“旋转角度 平移量”的组合。卷积的第一步对核 ψ 施加一个群变换 g^{-1}第二步与被变换后的输入 f(h) 做内积并对整个群 G 求和。看起来只是把求和范围从平面换成了群但这一步推广的意义非常深远它让卷积从“只共享平移参数”扩展到了“共享整个对称群的参数”。这里有一个容易绕晕的点第一层卷积的输入是普通图像定义在平面 Z² 上而群卷积要求输入定义在群 G 上。所以第一层卷积和后面的群卷积不是同一种操作。论文里把第一层叫做“升维卷积”lifting convolution它的作用是把平面上的信号“抬”到群上去(f₁ * ψ)(g) Σ_{x} f₀(x) ψ(g^{-1} x)直观理解就是输入图像只有一个“平面视角”但你同时用 4 个已经旋转过的核去卷积它每个旋转角度对应一个输出“视角”。这四个视角叠在一起就构成了群上的特征映射。从这之后信号就住在群里了后续层全部使用标准的群卷积。2.3 为什么群卷积一定等变等变性证明的直觉论文最漂亮的部分是用短短几行数学证明群卷积对群作用是等变的。设 L_g 表示对输入施加群变换 g那么有L_g(f₁ * ψ) (L_g f₁) * ψ这个公式的意思是先把输入旋转一下再做卷积和先做卷积再把输出旋转一下两者结果完全一样。这正是“等变”的定义。为什么能成立核心原因来自群乘法的结合律。卷积在本质上是在做“匹配”当你把输入整体做了一个变换时你只需要对核或求和变量做相应的逆变换所有项的内积结果就会精确地对应回去。我刚开始读的时候觉得这个证明太短了像变魔术。后来自己动手做了一个小实验才真正理解拿一张数字图像用 P4-升维卷积算一遍然后把输入旋转 90 度再算一遍输出特征映射恰好是原来那 4 个通道的循环移位。那一刻我才意识到等变性不是一个“近似性质”而是一个精确的代数恒等式。这种精确性是纯靠数据增强永远无法得到的。3. 网络架构解析3.1 升维层从平面信号到群信号升维层是 GCNN 和普通 CNN 的第一个分水岭。它的输入是普通图片张量 [B, C, H, W]输出是一个住在 P4 群上的特征张量 [B, C, 4, H, W]其中新增的那个维度大小等于群的旋转次数。这个操作具体怎么做以 P4 为例准备一组普通卷积核对每个核分别旋转 0 度、90 度、180 度、270 度得到 4 份旋转后的核拿这 4 份核分别对输入图像做标准二维卷积得到 4 份输出特征图最后把这 4 份特征图沿旋转维度堆叠起来。需要特别注意的是升维层的参数并不是“每个旋转角度一套参数”而是“只有一套基础参数通过旋转复用”。这一点保证了等变性成立也是参数效率提升的根源。3.2 群卷积层特征映射在群上的卷积升维之后每一层的特征映射都是定义在群上的。群卷积层要做的事情是对输入特征映射的每个“旋转面”和输出特征映射的每个“旋转面”用对应的旋转卷积核做二维卷积然后在旋转维度上求和。这听起来有点绕。让我用一个具体例子说明。假设输入特征有 C_in 个通道、4 个旋转方向输出特征有 C_out 个通道。我们需要一组“群核”它的形状是 [C_out, C_in, 4, kh, kw]。对于输出旋转方向 r_out 的某个通道需要遍历输入旋转方向 r_in找到对应的二维核对输入特征的第 r_in 个旋转面做空间卷积然后对 r_in 求和。这里的实现细节很像普通卷积中“通道维度求和”的扩展版本普通卷积是在输入通道上求和群卷积则是在“输入通道 × 输入旋转方向”上求和。由于旋转方向只有 4 个计算量只是普通卷积的 4 倍但参数量被共享机制大幅压缩了。3.3 群池化把等变特征变成不变特征网络提取了等变特征之后最终分类任务还是需要一个固定维度的输出。这时候就需要“群池化”group pooling操作。最简单的方式是在旋转维度上取最大值或平均值max-pooling 相当于“我相信某个角度最像这个类别”average-pooling 相当于“所有角度的响应都投票”。群池化通常放在网络最后紧跟着全局平均池化。它把 [B, C, 4, H, W] 压缩回 [B, C, H, W]彻底丢弃方向信息从而实现旋转不变性。这也是整个网络“前段等变、后端不变”的关键一步。我在实验中发现一个有意思的现象如果中间层不小心也加了群池化网络的辨别能力会明显下降。这说明方向信息在中间层是有用的过早丢弃方向信息反而有害。这也印证了论文里反复强调的中间层要的是等变不是不变。4. 实验设置与结果分析4.1 旋转 MNIST 上的核心实验论文里最经典的一组实验是旋转 MNIST。普通 MNIST 测试集上的手写数字都是正立的模型很好认但如果把测试集所有数字旋转任意角度普通 CNN 的性能会急剧下降。论文构造了一个训练集仍然接近正立、测试集包含随机旋转的设定用来检验模型对旋转的泛化能力。结果是压倒性的GCNN 在旋转测试集上的错误率远低于标准 CNN甚至在标准 CNN 加上旋转数据增强之后GCNN 仍然有明显优势。这验证了我前面讲的直觉数据增强只是在“见过”的角度上拟合而 GCNN 的结构保证让网络对“没见过”的角度也能做出正确响应。4.2 在 CIFAR-10 和 STL-10 上的验证除了 MNIST论文还在 CIFAR-10 和 STL-10 上做了实验。在这些更复杂、通道更多的数据集上GCNN 也展现出了稳定的提升。不过这里的提升幅度比旋转 MNIST 要小一些原因也好理解真实图片的旋转模式更复杂自然图像中的目标不是只有 90 度旋转这一种对称性而且 CIFAR-10 的训练数据量相对充足数据增强已经能覆盖不少旋转变化。但有一个趋势非常清晰数据量越少GCNN 相对普通 CNN 的优势越大。GCNN 因为权重共享机制天然携带了更强的正则化在小数据集上几乎不会过拟合而普通 CNN 很容易把各个旋转角度的特征重复学好几遍。4.3 参数效率和内禀正则化的来源为什么 GCNN 参数更少却效果更好关键在于“共享”的粒度不同。普通 CNN 的权重只在平移方向共享而 GCNN 的权重在平移和旋转两个方向上都共享。同一组边缘检测参数在 4 个旋转角度下复用网络就不需要为每个角度单独存储一套参数。这相当于给网络加了很强的先验我们相信图像结构在旋转下是被保住的。先验越强模型需要从数据中学习的自由度就越少自然在数据有限的情况下表现更好。当然如果数据量无限大、增强策略足够完备普通 CNN 也能逼近同样的性能但代价是参数和算力成倍增加。GCNN 提供的是“花小钱办大事”的杠杆。5. 代码实现与复现要点5.1 用 PyTorch 实现 P4 等变卷积我复现 GCNN 的时候把 P4 群卷积拆成了两个核心操作核旋转和群维度的循环卷积。下面这段代码展示升维层的实现思路import torch import torch.nn as nn import torch.nn.functional as F def rotate_kernel(kernel, k): 将卷积核逆时针旋转 90*k 度 kernel 形状: [out_c, in_c, kh, kw] if k 0: return kernel return torch.rot90(kernel, k, dims[2, 3]) class LiftConv(nn.Module): 将平面输入提升到 C4 群上的卷积层 def __init__(self, in_c, out_c, kernel_size3, group_size4): super().__init__() self.group_size group_size padding kernel_size // 2 self.conv nn.Conv2d( in_c, out_c, kernel_size, paddingpadding, biasFalse ) def forward(self, x): # x: [B, in_c, H, W] outs [] for k in range(self.group_size): kernel rotate_kernel(self.conv.weight, k) outs.append(F.conv2d(x, kernel, padding1)) # 堆叠后形状: [B, out_c, group_size, H, W] return torch.stack(outs, dim2)这里最关键的是torch.rot90的dims参数。dims[2, 3]表示在二维卷积核的后两个维度高度和宽度上旋转。我在第一次写的时候用成了dims[0, 1]结果把输出通道维度和输入通道维度拧在一起了模型直接不能收敛。这种维度错误非常隐蔽建议你在实现后立即做等变性测试下面会讲。5.2 群卷积层的实现思路群卷积层比升维层复杂一些因为要在群维度上做求和。我采用的是最直观的循环实现虽然效率不高但逻辑清晰适合作为复现起点class GroupConv(nn.Module): C4 上的群卷积层简化版实现 def __init__(self, in_c, out_c, kernel_size3, group_size4): super().__init__() self.group_size group_size padding kernel_size // 2 # 群核: 对每个输出旋转方向和每个输入旋转方向各有一份核 self.weight nn.Parameter( torch.randn(out_c, in_c, group_size, kernel_size, kernel_size) ) self.bias nn.Parameter(torch.zeros(out_c, group_size)) def forward(self, x): # x: [B, in_c, R, H, W], R group_size B, in_c, R, H, W x.shape out_c self.weight.shape[0] outputs [] for r_out in range(R): acc 0 for r_in in range(R): # 核心技巧输入的第 r_in 面对应核旋转 (r_out - r_in) k (r_out - r_in) % R kernel rotate_kernel(self.weight[:, :, r_in], k) acc F.conv2d( x[:, :, r_in], kernel, padding1 ) outputs.append(acc) out torch.stack(outputs, dim2) out out self.bias.view(1, 1, R, 1, 1) return out这版实现用了两层循环在 group_size 只有 4 的时候速度还能接受但扩展到 P4M8 个群元素或 P88 个旋转方向时就很慢了。实际工程中更高效的做法是用group convolution或通道重排来并行化但那些优化代码会掩盖数学结构不利于理解论文。我在复现时先用这个简单版本跑通逻辑再用torch.einsum优化。5.3 等变性自检一个必须写的测试复现 GCNN 最容易犯的错误是“代码看起来没问题但等变性根本不成立”。为了避免这种情况我强烈建议你在写完每一层之后都加一个等变性测试def test_equivariance(model, x, group_size4): 验证模型对 C4 旋转是否等变 # 计算原始输出的旋转后结果 out1 model(x) # [B, C, R, H, W] out1_rot torch.stack( [out1[:, :, (i - k) % group_size] for k in range(group_size)] ) # 计算旋转输入的输出 x_rot torch.rot90(x, 1, dims[2, 3]) out2 model(x_rot) # 比较 diff (out1_rot - out2).abs().max().item() assert diff 1e-4, f等变性测试失败, 最大差异: {diff}这个测试的原理是输入旋转 90 度后群卷积的输出应该等价于原始输出的群维度循环移位。如果你实现了 P4M 群还需要额外考虑反射维度的对应关系。我用这个方法抓到过一个很隐蔽的 bug因为torch.rot90对dims顺序敏感导致核的旋转方向和预期符号相反模型最后学成了一堆没意义的特征。5.4 训练配置与实用建议复现时我用 Adam 优化器初始学习率 1e-3batch size 64在旋转 MNIST 上手写数字分类任务大概 60 到 80 个 epoch 就能收敛。实际上由于 GCNN 的正则化强度更大它的收敛速度比普通 CNN 略慢但最终精度更高。我建议把学习率调低一个数量级试试因为共享参数的等变特征对梯度步长更敏感。还有一个容易被忽略的细节输入归一化。GCNN 对输入做了旋转增强后像素均值保持不变但不同旋转方向的方差分布有细微差异。我在实验中发现先在普通图像上计算每个通道的 mean 和 std再对旋转后的图像应用同样的归一化参数效果比从小到大逐层归一化好得多。另外提醒一下如果你要处理的是非方形输入比如 32x48 的图像旋转 90 度会改变宽高比这就不仅仅是群维度上的循环移位了。论文里的实验基本都用方形输入。实际项目中遇到这种情况要么 resize 成方形要么改用只含镜像不含旋转的群比如把 P4 换成 D1总之要保证变换本身不自相矛盾。6. 对后续工作的影响与个人思考6.1 GCNN 打开了等变网络的一扇门GCNN 发表之后等变网络迅速成为一个活跃方向。后续工作大体沿两条线推进一条是位点等变网络Steerable CNN把旋转等变从离散群推广到连续旋转群并用球谐函数等数学工具构造等变卷积核另一条是通用等变框架比如把等变思想推广到图网络、点云、分子结构上。你会发现这些工作几乎都引用 GCNN 作为起点因为它第一次明确给出了“把群的对称性嵌入卷积结构”的通用范式。我个人觉得 GCNN 最大的贡献不是某个具体的网络结构而是提供了一种思考方式当你发现网络对某种变换不够鲁棒时应该先想“这种变换有没有数学上的群结构”如果有就把它写成结构约束而不是靠数据增强疯狂堆样本。这种做法更优雅、更省资源而且泛化边界更清晰。6.2 GCNN 的局限与改进方向作为一篇 2016 年的工作GCNN 自然有局限。最明显的是它只处理离散旋转群对连续角度的旋转只能近似而且 P4 和 P4M 群的生成元有限不能覆盖所有可能的平面刚体变换。另外它的计算开销也是问题虽然参数少了但群卷积的中间特征图比普通 CNN 大了 R 倍显存和推理时间都随之上涨。我个人在使用中的体会是GCNN 最适合那些“旋转歧义明确、方向变化可枚举”的任务比如医学影像里细胞核方向的随机旋转、工业检测中零部件的多角度摆放。如果任务本身方向影响很小或者数据量大到数据增强已经完全覆盖了旋转分布那 GCNN 的性价比就没那么高了普通 CNN 加增强可能更省事。选不选它本质上是在“更强的先验”和“更灵活的特征表达”之间做权衡。
返回列表