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×224,patch 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_channels=3, embed_dim=96, patch_size=4): super().__init__() self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_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_size+1, 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], indexing="ij")) 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、HGFormer)
Swin 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.7+,GPU显存至少16G,我使用的是单张RTX 4090。如果显存不够,可以用Swin-Tiny作为backbone,batch 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/8,stage3输出1/16,stage4输出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卷积负责细化。反卷积容易产生棋盘格伪影,在深度图上表现为细密纹理条纹,双线性插值加卷积效果更稳。
输出层不接激活函数,因为深度值可以大于1,sigmoid或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-4,weight 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骨干。这些方向我都做过初步尝试,踩坑不少但收获更多,后面有机会再单独写文章展开。