多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

Transformer连续像素级预测:从自注意力到深度估计

Transformer连续像素级预测:从自注意力到深度估计 1. 项目解剖为什么Transformer能做连续像素级预测1.1 这个标题到底在解决什么问题先把这个标题拆开看。Transformer-Based Attention Networks指以自注意力为核心的Transformer架构Continuous Pixel-Wise Prediction指的是对图像每一个像素输出一个连续数值而不是类别标签。放在一起就是一套用注意力机制做密集回归任务的完整技术路线。这个路线覆盖的任务远比想象中广。单目深度估计是每个像素预测一个距离值光流估计是每个像素预测一个二维位移向量表面法线估计是每个像素预测一个三维方向还有密度图预测、散焦估计、视差估计等等。这些任务有一个共同特点输出和输入分辨率一致且每个位置的预测值是连续实数不是离散类别。这和图像分类、目标检测、语义分割这类任务有本质区别分类任务最后接一个softmax就行而连续像素级预测需要的是回归头输出层连的是激活函数或者干脆不连激活函数损失函数也完全是另一套玩法。1.2 为什么不用纯CNN而要用Attention很多人一开始会问CNN做深度估计都做了快十年了从Eigen的multi-scale网络到DORN、BTS效果也一直在涨为什么非要换Transformer我的理解是CNN受限于局部感受野。虽然可以通过堆叠卷积层、扩大卷积核、使用空洞卷积来增大感受野但本质上卷积算子建模的是局部邻域的加权求和全局依赖关系需要靠很多层去“传递”。对于深度估计这类需要全局理解的任务来说这存在先天短板。举个例子一张图中地面远处有一辆车要估计这辆车的深度算法必须理解“地面是连续平面”“车和地面的空间关系”“远处物体整体尺度缩小”这些全局信息。CNN在浅层看到的是局部纹理只有层层上采样、不断汇聚上下文之后才能形成全局判断这个过程效率低而且容易在长距离依赖上出现信息丢失。注意力机制完全不同。自注意力一步到位把任意两个位置之间的距离缩短为一次矩阵运算不管两个像素在图像上相隔多远注意力权重都能直接建立联系。这种全局感受野让Transformer在理解场景结构、物体遮挡关系、空间连续性等方面天然占优。对于连续像素级预测来说好的全局理解意味着深度边界更清晰、平面区域更平滑、物体相对位置更合理。1.3 连续像素级预测的技术挑战把Transformer搬到像素级预测上不是简单地把分类头换成回归头就完事有四个绕不开的问题。第一计算复杂度。标准自注意力的复杂度是序列长度的平方一张512×512的图切成16×16的patch序列长度是1024注意力矩阵是1024×1024计算量尚可接受但切到8×8的patch序列长度变成4096注意力矩阵就是4096×4096显存直接爆炸。像素级预测需要保留空间细节通常需要高分辨率输入这就和标准Transformer的平方复杂度正面冲突。第二多尺度问题。连续像素级预测任务中物体大小差异极大。近处的一辆车可能占据几百个像素远处的一辆车只有几十个像素。Transformer虽然能捕捉全局依赖但如果只在单一尺度上做自注意力小物体细节会丢失。要让模型在预测深度或光流时对小物体和大结构都能兼顾必须有特征金字塔或层级化结构。第三边缘模糊问题。注意力机制擅长捕捉全局结构但密集回归任务对局部边界非常敏感。很多基于CNN的方法已经能预测出比较清晰的边缘而第一代基于纯Transformer的方法经常出现深度图整体平滑但边界模糊的情况原因就在于patch化操作patch embedding本身丢掉了部分像素级细节。第四训练难度。Transformer比CNN更依赖大数据优化起来也更敏感。深度估计数据集比如NYU Depth V2、KITTI样本量比ImageNet小一两个数量级直接用原始ViT的默认超参数从头训练收敛慢且容易过拟合。所以实操中需要做大量的训练策略适配。2. 架构思路拆解从全局感知到密集输出2.1 从patch embedding开始图像如何变成序列Transformer最标准的图像输入方式是ViT提出的patch embedding。把一张H×W×3的图像切成P×P的小块每个小块展平成向量再用一个线性层映射到D维。假设输入是224×224patch size是16那么序列长度是(224/16)×(224/16)196每个token的维度是D。选择patch size是个权衡。patch越小序列越长计算量越大但空间细节保留得越好。对于连续像素级预测来说16×16的patch通常太粗糙了尤其是输出需要恢复到原始分辨率时16倍下采样的细节损失很难弥补。实际应用中我倾向于先让浅层保持较高的分辨率再逐步下采样这也是为什么层级化Transformer比如Swin在密集预测任务中表现优于原始ViT的核心原因。代码层面patch embedding可以用一个卷积层实现卷积核大小和步长都等于patch size这样既完成了切块又完成了线性映射一步到位import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels3, embed_dim96, patch_size4): super().__init__() self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) self.norm nn.LayerNorm(embed_dim) def forward(self, x): # x: [B, 3, H, W] x self.proj(x) # [B, embed_dim, H/patch, W/patch] B, C, H, W x.shape x x.flatten(2).transpose(1, 2) # [B, H*W, embed_dim] x self.norm(x) return x这里注意一下LayerNorm放在卷积之后是Swin系列的标准做法先做归一化再进入Transformer block训练更稳定。reshape时先flatten再transpose得到的序列顺序就是从左到右、从上到下符合空间位置的自然排列。2.2 位置编码里的门道Transformer本身没有顺序概念自注意力是置换等变的所以要靠位置编码给token注入空间位置信息。ViT用的是绝对位置编码把每个位置学到一个固定向量加到patch embedding上。Swin用的是相对位置偏置在计算注意力矩阵时根据两个token之间的相对位置查表得到一个偏置项加在注意力分数上。对于像素级预测我强烈推荐相对位置偏置。原因很直接绝对位置编码对平移不敏感模型必须从大量数据中学习“位置A和位置B的相对关系”而相对位置编码直接把这个关系显式编码了让attention更容易学到相邻像素有高相关性、远处像素有低相关性这一先验。在深度估计和光流估计这种天然依赖空间连续性的任务上相对位置编码能明显加速收敛。Swin里的相对位置偏置实现并不复杂核心是维护一个可学习的偏置表然后用相对坐标索引去查表class RelativePositionBias(nn.Module): def __init__(self, num_heads, window_size): super().__init__() self.window_size window_size # 相对位置坐标范围是 [-window_size1, window_size-1] self.bias_table nn.Parameter( torch.zeros((2 * window_size - 1) * (2 * window_size - 1), num_heads) ) self.register_buffer(relative_position_index, self._get_relative_position_index()) def _get_relative_position_index(self): coords torch.arange(self.window_size) coords torch.stack(torch.meshgrid([coords, coords], indexingij)) coords coords.flatten(1) relative_coords coords[:, :, None] - coords[:, None, :] relative_coords relative_coords.permute(1, 2, 0).contiguous() relative_coords[:, :, 0] self.window_size - 1 relative_coords[:, :, 1] self.window_size - 1 relative_coords[:, :, 0] * 2 * self.window_size - 1 index relative_coords.sum(-1) return index def forward(self, q_size): # 返回 [num_windows*num_heads, seq_len, seq_len] 的偏置 return self.bias_table[self.relative_position_index].permute(2, 0, 1)2.3 自注意力机制的三个核心计算自注意力在一次前向过程中完成三件事算相关性、归一化、加权聚合。具体来说每个token生成三个向量Query表示“我想找什么”Key表示“我是什么”Value表示“我提供什么信息”。Query和Key做点积得到注意力分数经过softmax归一化后作为权重对Value加权求和。公式很简单Attention(Q,K,V) softmax(QKᵀ/√d) V但真正影响效果的是几个细节。除以√d是为了防止点积结果过大导致softmax进入饱和区梯度消失。多头注意力则是把D维空间分成h个子空间每个头独立做attention最后拼接起来。多个头的好处是每个头可以关注不同类型的依赖关系有的头看颜色相似性有的头看位置邻近性有的头看纹理一致性合在一起表达能力更强。一个容易忽略的点是在像素级预测任务中注意力矩阵本身蕴含了丰富的空间关系信息。比如在深度估计中注意力权重大的区域往往是位于同一平面或同一物体上的像素。有一些工作直接把注意力图作为特征传递给解码器效果比只用Value的加权结果更好。这个做法实现起来很简单就是把多头注意力输出的attention map做pooling或reshape后concat到特征里但收益明显。2.4 层级化设计的价值Swin、HGFormerSwin Transformer贡献了窗口注意力window attention和移位窗口shifted window让视觉Transformer第一次有了真正意义的层级化特征。图像先切成小窗口在每个窗口内部做自注意力窗口数量固定计算复杂度和图像尺寸呈线性关系而不是平方关系。通过交替使用规则窗口和移位窗口让信息能在窗口之间流动打破了窗口内的局部限制。Swin给出的下采样金字塔特征让密集预测任务可以像用ResNet那样直接套用FPN、UNet等成熟结构这非常关键。窗口注意力本质上是一个“局部先验”类似于卷积但比卷积更灵活这正是连续像素级预测需要的。HGFormer这类工作进一步往前走了一步把超图学习hypergraph learning引入Transformer。超图与普通图的区别在于普通图的一条边连接两个节点超图的一条超边可以连接任意数量的节点天然适合表达多像素之间的高阶关系。比如判断某个像素是否属于同一个物体表面不能只看两两关系可能需要同时看到一整块区域的像素才会更有把握。HGFormer用超图卷积补充自注意力的边关系建模让模型对拓扑结构更敏感。这类“结构感知”的设计对于深度估计中处理复杂遮挡、拓扑关系明显的场景非常有价值。3. 实操落地搭建一个可用于像素级预测的Transformer基线3.1 环境与数据准备我建议以单目深度估计作为切入点因为它最能体现连续像素级预测的特点而且数据集好找、评估指标直观。数据集用NYU Depth V2官方切分训练集约2.4万张测试集654张评估指标看绝对相对误差Abs Rel、均方根误差RMSE和δ1准确率。环境方面PyTorch 1.13或2.0以上CUDA 11.7GPU显存至少16G我使用的是单张RTX 4090。如果显存不够可以用Swin-Tiny作为backbonebatch size设为8输入分辨率降到320×240显存占用大概12G左右不影响实验验证。预处理要做三件事图像缩放到固定分辨率、随机水平翻转和随机颜色扰动、深度值归一化。深度估计有一个特殊之处NYU深度图有大量无效区域一般把深度大于10米的截断为10然后除以10归一化到0到1之间让回归目标处于一个相对合理的数值范围模型更好优化。3.2 编码器核心代码拆解这里我给一个简化版层级化Transformer编码器参考Swin的设计但砍掉了复杂部分保留窗口注意力、patch merging和层级特征输出便于从头理解。窗口注意力部分的核心是窗口划分和窗口还原。把特征图按window_size切块每块内部做attention处理完再拼回去def window_partition(x, window_size): B, H, W, C x.shape x x.view(B, H // window_size, window_size, W // window_size, window_size, C) x x.permute(0, 1, 3, 2, 4, 5).contiguous() x x.view(-1, window_size, window_size, C) return x def window_reverse(windows, window_size, H, W): B int(windows.shape[0] / (H * W / window_size / window_size)) x windows.view(B, H // window_size, W // window_size, window_size, window_size, -1) x x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) return xStage模块负责把输入从高分辨率映射到不同尺度。每个stage包含patch embedding同时改变分辨率和通道数和若干个window attention block。在像素级预测任务中stage1通常输出1/4分辨率特征stage2输出1/8stage3输出1/16stage4输出1/32这样接解码器时多尺度特征才完整。一个实操经验是stage1的窗口大小不要设太大。输入是320×240时1/4分辨率下特征图是80×60如果窗口设成8×8窗口内部只有64个token局部建模能力绰绰有余。到了stage3、stage4特征图缩小到1/16和1/32窗口内token数量已经很少这时窗口注意力更像是在做小块内部的全局建模信息瓶颈会比较明显。3.3 解码器与连续预测输出解码器负责把多尺度特征逐步恢复到原始分辨率并输出连续像素预测值。我采用的是类FPN结构加UNet风格的跳跃连接比直接用单层上采样效果好很多。具体做法stage4特征经过一个卷积层调整通道数后上采样2倍与stage3的特征concat再过一个3×3卷积和上采样依此类推逐步融合stage2和stage1的特征。最后经过一个1×1卷积输出单通道深度图。整个过程不需要反卷积上采样用双线性插值就行后面的3×3卷积负责细化。反卷积容易产生棋盘格伪影在深度图上表现为细密纹理条纹双线性插值加卷积效果更稳。输出层不接激活函数因为深度值可以大于1sigmoid或ReLU都会限制输出范围。深度估计的常见做法是输出log深度训练时用log空间的损失函数这样近距离误差和远距离误差在损失中占比更均衡避免模型过度优化近距离的大数值差异。3.4 训练策略与损失函数选择训练Transformer做回归任务损失函数的选择直接决定模型行为的“性格”。分类任务用交叉熵而连续像素级预测有多个常用选择L1损失对异常值不敏感收敛稳定但梯度在零点不连续。L2损失MSE对大误差惩罚更大但容易被离群点带偏。BerHu损失结合两者优点小误差时用L2大误差时用L1是深度估计的经典选择。尺度不变损失按像素对数误差计算忽略整体尺度偏移适合深度估计。我的经验是单目深度估计优先试BerHu光流估计优先试L1加平滑项。如果用了log深度输出可以把L1和梯度平滑损失组合起来让预测深度图在局部区域内更平滑但边缘处保留跳变。常见做法class DepthLoss(nn.Module): def __init__(self): super().__init__() self.criterion nn.SmoothL1Loss() def forward(self, pred, gt, mask): pred pred[mask] gt gt[mask] loss self.criterion(pred, gt) # 梯度平滑项对预测图的水平和垂直方向梯度做L1约束 dx torch.abs(pred[:, :, 1:, :] - pred[:, :, :-1, :]).mean() dy torch.abs(pred[:, :, :, 1:] - pred[:, :, :, :-1]).mean() loss loss 0.1 * (dx dy) return loss训练参数上Transformer比CNN更挑剔。优化器用AdamW初始学习率2e-4weight decay设为0.05线性warmup 1000步之后cosine decay。batch size在单卡上能开到多大就多大实验证明Transformer在小batch size下收敛明显变慢。如果显存有限宁可降低分辨率也不要过度减小batch size我实测batch size从8降到4同样训练轮数指标会掉5%到8%。4. 常见问题与排查技巧实录4.1 显存爆炸与分辨率限制怎么破这是把Transformer接到密集预测任务上最先遇到的问题。标准ViT处理高分辨率图像时显存快速增长解决办法按优先级排列第一优先是改用层级化结构Swin或PVT让高分辨率阶段只在浅层做窗口内注意力第二是使用窗口注意力把全局attention改成局部attention显存从平方级降到线性级第三是开启梯度检查点gradient checkpointing用计算换显存训练速度大约下降20%到30%但显存需求能降低一半。我在1280×720输入的情况下用梯度检查点加窗口注意力成功把batch size从2提到6。4.2 预测结果边缘模糊、细节丢失最典型的问题是深度图整体结构没问题但物体边缘糊成一片。我排查了三个点第一个是patch size使用4×4甚至2×2的patch作为第一阶段embedding保留更多空间细节第二个是解码器结构单纯从1/32分辨率一路upsample到原图必然丢失边缘需要跨尺度跳跃连接把浅层高分辨率特征反复融入第三个是损失函数只靠全局L1损失会让模型倾向于输出平滑结果可以补充一个边缘感知损失用预测深度的梯度与输入图像梯度的相似性作为额外约束。4.3 模型收敛慢或训练震荡Transformer训练不稳定的原因大多是学习率策略和初始化不匹配。我试过直接用较大的恒定学习率训练loss曲线像心电图一样抖动换上warmup加cosine decay之后明显稳定。还有一个经常踩坑的地方attention block内部的LayerNorm位置。Pre-LN结构norm放在attention之前比Post-LN结构训练稳定得多建议所有Transformer block都用Pre-LN。如果震荡还是严重检查一下是否忘了在patch embedding后面加LayerNorm这个位置缺失会导致深层特征分布漂移。4.4 拓扑结构复杂场景效果差复杂场景下比如多物体互相遮挡、树冠间隙、透明物体基于局部窗口的attention经常建模失败。HGFormer的思路可以参考把图像特征构建成超图一个超边连接多个被判定为同一语义区域的像素再做超图卷积更新节点特征。超图的好处是能一次性建立多点之间的高阶约束尤其适合表达“多个像素共同属于一个物体平面”这种非两两关系。实操中先用自监督方式把特征聚成簇把同一簇的像素作为一条超边然后在这些超边上做消息传递。这个思路可以用在解码器部分在FPN的最后一层输出前加一个超图卷积模块既控制了整体算力开销又能明显改善拓扑复杂区域的预测质量。4.5 一个容易被忽视的细节深度归一化与逆深度NYU和KITTI这类数据集的深度标签范围差异很大直接把原始深度值喂给模型会让损失被远距离样本主导。除了截断归一化还有一个更激进的做法预测逆深度1/depth。逆深度在自动驾驶场景中更符合成像几何近处物体的深度精度更高模型对近距离障碍物更敏感。实际对比过两种目标表示逆深度在KITTI的Abs Rel指标上能提升2个点左右但近距离噪声会被放大需要配合更平滑的损失项。最后分享一点实际体会把Transformer用到连续像素级预测上最关键的思维转变是不要把它当成一个可以即插即用的黑盒而是要理解它和CNN在归纳偏置上的本质差异。CNN把局部连续性当作硬先验写死在架构里Transformer把一切关系都交给注意力去学习前者在数据少时更容易收敛后者在数据足时上限更高。实际项目中手工调一个ResNet加空洞卷积的基线可能只需要一天而Transformer方案从设计到调通至少需要一两周但一旦跑通在复杂场景下的泛化能力通常会让之前的CNN方案难以企及。另外如果想在这个方向做进一步探索可以沿着三个方向走一是把注意力图和超图结构可视化出来观察模型在不同场景下关注哪些区域这对调试非常有帮助二是尝试轻量化方案用蒸馏或剪枝把大模型压缩到能跑在边缘设备上三是把连续像素级预测扩展成多任务统一框架让深度估计、表面法线估计、语义分割共享一个Transformer骨干。这些方向我都做过初步尝试踩坑不少但收获更多后面有机会再单独写文章展开。
返回列表