共计 1489 个字符,预计需要花费 4 分钟才能阅读完成。
1. 背景:传统图像修复方法的局限性
图像修复任务旨在从损坏或退化的输入中恢复高质量图像内容。传统方法主要面临两个核心挑战:
-
CNN 的局部性限制 :传统卷积神经网络(CNN)受限于局部感受野,难以建模长距离像素依赖关系。当处理大面积破损(如遮挡区域)时,修复结果常出现结构扭曲或纹理模糊。
-
普通 Transformer 的计算瓶颈 :Vision Transformer 虽然能建模全局关系,但其自注意力计算复杂度随图像分辨率呈平方级增长。对于 512×512 图像,标准 Transformer 的 FLOPs 高达 $O((512×512)^2)$,导致训练和推理成本剧增。
2. Attention Retractable Transformer 核心技术解析
2.1 层级收缩机制设计

图:通过收缩因子 k 动态调整注意力范围
核心创新点在于引入可学习的收缩因子 $k$,将原始 $H×W$ 特征图划分为 $\frac{H}{k}×\frac{W}{k}$ 的窗口,在窗口内计算局部注意力:
$$
\text{Attention}(Q,K,V) = \text{Softmax}(\frac{QK^T}{\sqrt{d_k}} + M)V
$$
其中 $M$ 为动态生成的掩码矩阵,控制不同区域的注意力范围。
2.2 计算效率对比
| 模型 | FLOPs (512×512) | 显存占用 |
|---|---|---|
| Swin-T (window=8) | 28.3G | 6.2GB |
| ART (k=4) | 17.1G | 3.8GB |
| ART (动态 k) | 12.4G~19.7G | 3.2GB |
3. PyTorch 实现关键代码
class MultiHeadRetractableAttention(nn.Module):
def __init__(self, dim, num_heads, retract_ratio=4):
super().__init__()
self.scale = (dim // num_heads) ** -0.5
self.qkv = nn.Linear(dim, dim * 3)
self.proj = nn.Linear(dim, dim)
# 可学习收缩因子(初始化为 ratio)self.retract_k = nn.Parameter(torch.tensor(float(retract_ratio)))
def forward(self, x, H, W):
B, N, C = x.shape
k = round(self.retract_k.item()) # 动态取整
# 重采样到低分辨率空间(关键优化)x_low = F.adaptive_avg_pool2d(x.view(B,H,W,C).permute(0,3,1,2), (H//k, W//k))
qkv = self.qkv(x_low.flatten(2).transpose(1,2))
# ... 后续注意力计算...
4. 实验验证
4.1 定量结果
在 CelebA-HQ 测试集上(掩码率 30%~50%):
| Method | PSNR↑ | SSIM↑ | Time↓ |
|---|---|---|---|
| ContextAE | 28.7 | 0.891 | 0.23s |
| SwinIR | 30.1 | 0.907 | 0.41s |
| ART (ours) | 31.4 | 0.923 | 0.29s |
4.2 注意力可视化
不同收缩因子下的注意力分布(k=1/2/4)
5. 生产部署建议
- 动态收缩策略 :
- 对高频纹理区域采用小 k 值(精细修复)
-
对平滑背景区域采用大 k 值(节省计算)
-
混合精度训练 :
- 对收缩因子 k 采用 FP32 保持精度
-
其他参数使用 FP16 加速
-
边缘设备优化 :
- 使用 TensorRT 合并 qkv 投影和注意力计算
- 对重采样操作启用 CUDA Graph
该方案在 NVIDIA Jetson AGX 上实现 17ms 的单帧处理速度,满足实时修复需求。未来可探索与扩散模型的结合,进一步提升生成质量。
正文完
发表至: 计算机视觉
近一天内
