首页
/ Transformers 中 Swin2SR 模型详解:基于 Swin Transformer V2 的图像超分辨率与修复实现

Transformers 中 Swin2SR 模型详解:基于 Swin Transformer V2 的图像超分辨率与修复实现

2026-09-07 17:20:01作者:范垣楠Rhoda

本文围绕 Transformers 仓库中 Swin2SR 模型的官方文档展开,完整覆盖其设计动机、五大核心 API(Swin2SRConfigSwin2SRModelSwin2SRForImageSuperResolution 及两套图像处理器)的配置参数、前向流程与端到端用法,并结合 modeling_swin2sr.pyconfiguration_swin2sr.py 等源码剖析窗口注意力、上采样头与像素归一化等关键实现,帮助读者在仓库层面完整掌握该模型的结构与调用方式。

模型概述:从 SwinIR 到 Swin2SR

Swin2SR 由 Marcos V. Conde、Ui-Jin Choi、Maxime Burchi、Radu Timofte 提出(论文 Swin2SR: SwinV2 Transformer for Compressed Image Super-Resolution and Restoration,2022 年 9 月发布),于 2022 年 12 月由社区贡献者 nielsr 合入 Transformers(原始 PyTorch 代码来自 mv-lab/swin2sr 项目,见官方文档 docs/source/en/model_doc/swin2sr.md)。

其核心改进在于:将 SwinIR 中的 Swin Transformer 主干替换为 Swin Transformer V2 层,以此缓解 Transformer 视觉模型训练中的三类典型问题:

  1. 训练不稳定(training instability)——Swin V2 采用带可学习 logit scale 的余弦注意力;
  2. 预训练与微调之间的分辨率差距(resolution gap)——V2 的连续相对位置偏置(continuous relative position bias)使位置编码在窗口大小变化时更可迁移;
  3. 数据饥饿(hunger on data)——更高效的建模降低了对大规模数据的依赖。

论文在三个代表性任务上验证了该方法:JPEG 压缩伪影去除、图像超分辨率(经典与轻量级)以及压缩图像超分辨率,并曾位列 AIM 2022 压缩图像视频超分辨率挑战赛前五。

核心 API 一览

官方文档为以下五个类提供了 [[autodoc]] 接口文档,它们在仓库中的实现位置如下:

角色 实现文件
Swin2SRConfig 模型配置,定义网络深度、窗口大小、上采样因子等 configuration_swin2sr.py
Swin2SRModel 基础 Transformer 编码器(无任务头) modeling_swin2sr.py
Swin2SRForImageSuperResolution 带超分/修复上采样头的完整模型 modeling_swin2sr.py
Swin2SRImageProcessor Torchvision 后端图像处理器(rescale + 整除填充) image_processing_swin2sr.py
Swin2SRImageProcessorPil PIL/NumPy 后端图像处理器 image_processing_pil_swin2sr.py

Swin2SRConfig:配置参数详解

Swin2SRConfig 继承自 PreTrainedConfigmodel_type = "swin2sr"。从源码 configuration_swin2sr.py 可以看到全部参数及其默认值:

参数 默认值 说明
image_size 64 输入图像尺寸(标量或 [H, W]
patch_size 1 图像分块的 patch 大小
num_channels 3 输入通道数(RGB 为 3)
num_channels_out None(回退为 num_channels 输出通道数;__post_init__ 中若为 None 则置为 num_channels
embed_dim 180 Transformer 嵌入维度(即 hidden_size 别名)
depths [6, 6, 6, 6, 6, 6] 每个 RSTB 阶段的层数;num_layerslen(depths) 自动推导
num_heads [6, 6, 6, 6, 6, 6] 每个阶段的注意力头数(num_attention_heads 别名)
window_size 8 窗口注意力大小
mlp_ratio 2.0 前馈层扩展比
qkv_bias True Q/V 线性层是否带偏置(K 无偏置)
hidden_dropout_prob / attention_probs_dropout_prob 0.0 隐藏层与注意力概率 dropout
drop_path_rate 0.1 随机深度(Stochastic Depth)上限,按层线性分配
hidden_act "gelu" 前馈层激活函数
use_absolute_embeddings False 是否使用绝对位置嵌入(Swin 系默认用窗口相对位置编码)
initializer_range 0.02 截断正态初始化标准差
layer_norm_eps 1e-5 LayerNorm 的 epsilon
upscale 2 上采样因子:超分任务取 2/3/4/8;去噪与 JPEG 伪影去除取 1
img_range 1.0 输入图像数值范围,用于前向归一化
resi_connection "1conv" 每个 RSTB 阶段残差连接前使用的卷积块("1conv""3conv"
upsampler "pixelshuffle" 重建模块:'pixelshuffle'/'pixelshuffledirect'/'nearest+conv'/None

官方给出的最小化初始化示例(来自配置类 docstring):

from transformers import Swin2SRConfig, Swin2SRModel

# 初始化一个 caidas/swin2sr-classicalsr-x2-64 风格的配置
configuration = Swin2SRConfig()

# 基于该配置初始化模型(随机权重)
model = Swin2SRModel(configuration)

# 访问模型配置
configuration = model.config

另外注意 attribute_maphidden_size/num_attention_heads/num_hidden_layers 分别映射到 embed_dim/num_heads/num_layers,即通用接口读取这些标准属性时能拿到 Swin2SR 对应的私有字段。

架构剖析:Swin2SRModel 的前向流程

Swin2SRModel 是编码器主体,main_input_namepixel_values,支持梯度检查点(supports_gradient_checkpointing = True)。其结构(见 modeling_swin2sr.py)依次为:

  1. 图像级归一化参数:当输入/输出均为 3 通道 RGB 时,使用预计算的通道均值 [0.4488, 0.4371, 0.4040](非 3 通道时为全零),以 nn.Buffer 形式挂在 self.mean 上;
  2. first_convolution3x3 卷积将 num_channels 通道映射到 embed_dim
  3. Swin2SREmbeddings:内部 Swin2SRPatchEmbeddings 通过 stride=patch_size 的卷积做 patch 化并 LayerNorm,返回 patch 序列与特征图尺寸 input_dimensions
  4. Swin2SREncoder:按 depths 构建的 Swin2SRStage 列表,DropPath 概率通过 torch.linspace(0, drop_path_rate, sum(depths)) 在所有层上线性递增分配;
  5. 收尾LayerNormpatch_unembed 还原空间维度 → conv_after_body3x3 卷积),并与第一步卷积输出做残差相加。

forward 的关键预处理在 pad_and_normalize 中完成(modeling_swin2sr.py):

def pad_and_normalize(self, pixel_values):
    _, _, height, width = pixel_values.size()
    # 1. 用 reflect 填充到 window_size 的整数倍
    window_size = self.config.window_size
    modulo_pad_height = (window_size - height % window_size) % window_size
    modulo_pad_width = (window_size - width % window_size) % window_size
    pixel_values = nn.functional.pad(pixel_values, (0, modulo_pad_width, 0, modulo_pad_height), "reflect")
    # 2. 归一化
    mean = self.mean.to(device=pixel_values.device, dtype=pixel_values.dtype)
    pixel_values = (pixel_values - mean) * self.img_range
    return pixel_values

可以看到模型内部独立地做了 reflect 填充和 (x - mean) * img_range 归一化——这与图像处理器层面的预处理相互独立,推理时两者叠加:处理器负责 rescale 到 [0, 1] 与对称填充,模型负责窗口对齐与均值缩放。

RSTB 阶段与 Swin V2 注意力

Swin2SRStage 对应原始实现中的 Residual Swin Transformer Block(RSTB),每阶段内部:

  • 偶数索引层 shift_size = 0(窗口内注意力 MSA),奇数索引层 shift_size = window_size // 2(循环移位窗口注意力 SW-MSA),即经典的 MSA/SW-MSA 交替结构;
  • 阶段末尾按 resi_connection 配置插入 "1conv"(单个 3x3 卷积)或 "3conv"(两个 1x1 加一个 3x3 的窄通道序列,节省参数与显存),再与阶段输入残差相加(modeling_swin2sr.py)。

Swin2SRLayer 的前向流程(modeling_swin2sr.py)为:按 window_size 零填充 → torch.roll 做循环移位(仅移位层)→ window_partition 切窗 → 计算 3x3 区域划分的注意力掩码(get_attn_mask 用一次向量化区域计算替代了原始实现的 9 层 Python 循环)→ 注意力 → window_reverse 并反向移位 → 裁剪填充 → LayerNorm 残差 → FFN。

值得注意的两处 Swin V2 特征体现在 Swin2SRSelfAttention 中(modeling_swin2sr.py):

  • 带可学习 logit scale 的余弦注意力logit_scale 为可学习参数,qkF.normalize 再相乘并乘以 clamp(logit_scale, max=log(1/0.01)).exp(),这正是 Swin V2 提升训练稳定性的关键设计;
  • 连续相对位置偏置:不再查离散偏置表,而是用一个 Linear(2,512) -> ReLU -> Linear(512, num_heads) 的小 MLP,从归一化到 [-8, 8] 的相对坐标连续生成位置偏置,再经 16 * sigmoid(·) 缩放后加入注意力分数,从而支持预训练/微调窗口不一致的迁移场景。

Swin2SRForImageSuperResolution:上采样头与四种重建模块

Swin2SRForImageSuperResolutionSwin2SRModel 之上按 config.upsampler 挂载不同的重建头(modeling_swin2sr.py):

upsampler 取值 实现类 适用场景
"pixelshuffle" PixelShuffleUpsampler 经典超分:Conv → LeakyReLU → N 级 Upsample(2 的幂用逐级 Conv + PixelShuffle(2),3 用单次 PixelShuffle(3))→ 最终卷积
"pixelshuffledirect" UpsampleOneStep 轻量级超分:仅一步 Conv + PixelShuffle(scale),省参数
"nearest+conv" NearestConvUpsampler 真实世界超分(伪影更少);目前仅支持 upscale == 4,用两次 interpolate(scale_factor=2, mode="nearest") 加卷积堆实现
其他(如 None 单个 final_convolution 去噪 / JPEG 伪影去除:输出为 pixel_values + final_convolution(seq) 的残差形式

此外还支持 pixelshuffle_aux 分支(PixelShuffleAuxUpsampler),它会先用双三次插值把输入放大,经辅助卷积后与主干上采样结果相加,同时输出一个反归一化的 aux 张量。

前向流程(modeling_swin2sr.py)要点:

outputs = self.swin2sr(pixel_values, ...)
sequence_output = outputs[0]

if self.upsampler in ["pixelshuffle", "pixelshuffledirect", "nearest+conv"]:
    reconstruction = self.upsample(sequence_output)
elif self.upsampler == "pixelshuffle_aux":
    reconstruction, aux = self.upsample(sequence_output, bicubic, height, width)
    ...
else:
    # 去噪 / JPEG 伪影去除
    reconstruction = pixel_values + self.final_convolution(sequence_output)

# 反归一化 + 裁剪掉模型内部 reflect 填充产生的多余像素
reconstruction = reconstruction / self.swin2sr.img_range + self.swin2sr.mean
reconstruction = reconstruction[:, :, : height * self.upscale, : width * self.upscale]

输出为 ImageSuperResolutionOutput,其中 reconstruction 即最终高分辨率图像。需要特别注意的是 docstring 中明确的限制:当前不支持训练——forward 中一旦传入 labels 会抛出 NotImplementedError("Training is not supported at the moment"),该模型在 Transformers 中仅提供推理能力。

图像预处理:rescale 与整除填充

两套图像处理器(Torchvision 后端 Swin2SRImageProcessor 与 PIL 后端 Swin2SRImageProcessorPil)共享相同的核心行为:

  • do_rescale = Truerescale_factor = 1 / 255,将 uint8 图像缩放到 [0, 1]
  • do_pad = Truesize_divisor = 8(默认窗口大小),通过**对称填充(symmetric padding)**使高宽成为 8 的整数倍:
def pad(self, images, pad_size, size_divisor=8, **kwargs):
    """Pad images to make height and width divisible by size_divisor using symmetric padding."""
    height, width = images.shape[-2:]
    pad_height = (height // size_divisor + 1) * size_divisor - height
    pad_width = (width // size_divisor + 1) * size_divisor - width
    return tvF.pad(images, (0, 0, pad_width, pad_height), padding_mode="symmetric")

size_divisor 是文档中列出的 preprocess 专用参数(见 Swin2SRImageProcessorKwargs),可调用 processor(image, size_divisor=16, return_tensors="pt") 覆盖默认值 8,使其与模型 window_size 对齐。另外,构造函数仍接受旧的 pad_size 参数并自动映射到 size_divisor,保证旧版预处理器配置向下兼容。

Torchvision 后端的 _preprocess 还额外做了 group_images_by_shape 分组处理,同形状图像堆叠成批处理后再还原顺序;PIL 后端则逐张处理。

端到端推理示例

以下示例取自 Swin2SRForImageSuperResolution.forward 的官方 docstring(modeling_swin2sr.py),展示了经典 2x 超分模型的完整调用链:

import torch
import numpy as np
from PIL import Image
import httpx
from io import BytesIO

from transformers import AutoImageProcessor, Swin2SRForImageSuperResolution

processor = AutoImageProcessor.from_pretrained("caidas/swin2SR-classical-sr-x2-64")
model = Swin2SRForImageSuperResolution.from_pretrained("caidas/swin2SR-classical-sr-x2-64")

url = "https://huggingface.co/spaces/jjourney1125/swin2sr/resolve/main/samples/butterfly.jpg"
with httpx.stream("GET", url) as response:
    image = Image.open(BytesIO(response.read()))
# 准备模型输入
inputs = processor(image, return_tensors="pt")

# 前向传播
with torch.no_grad():
    outputs = model(**inputs)

output = outputs.reconstruction.data.squeeze().float().cpu().clamp_(0, 1).numpy()
output = np.moveaxis(output, source=0, destination=-1)
output = (output * 255.0).round().astype(np.uint8)  # float32 转 uint8
# 可用 Image.fromarray(output) 可视化

流程对应关系:AutoImageProcessor 按上文规则完成 rescale/填充 → 模型 pad_and_normalize 做窗口对齐与均值缩放 → 主干编码 → 上采样头重建 → 反归一化并裁剪回 height * upscale × width * upscale → 得到 reconstruction。官方文档(docs/source/en/model_doc/swin2sr.md)还提到,社区维护了 Swin2SR 演示 Notebook 集合(NielsRogge/Transformers-Tutorials 的 Swin2SR 目录)与图像超分 Demo Space(jjourney1125/swin2sr),可作为上手参考。

权重转换脚本与测试验证

仓库中还提供了配套工具与测试,可进一步印证上述实现:

  • 权重转换convert_swin2sr_original_to_pytorch.py 负责把 mv-lab 官方 PyTorch 检查点权重映射进 Transformers 格式,是 caidas/swin2SR-classical-sr-x2-64 等社区检查点入库的基础;
  • 模型测试tests/models/swin2sr/test_modeling_swin2sr.py 中的 Swin2SRModelTesterimage_size=32depths=[1, 2, 1]window_size=2upscale=2 等迷你配置实例化 Swin2SRConfigSwin2SRModel,并在 Swin2SRModelTester 中校验输入张量形状为 [batch, num_channels, image_size, image_size];测试中还引用了 Swin2SRImageProcessorPil 做端到端输入验证;
  • 图像处理器测试tests/models/swin2sr/test_image_processing_swin2sr.py 专门覆盖 rescale、对称填充及 size_divisor 行为。

从测试与源码结构看,Swin2SR 在 Transformers 中的定位是一个推理优先的图像复原模型:配置侧参数齐全、与 Swin/SwinV2 共享窗口切分与相对位置机制,工程侧则通过双后端图像处理器与独立 pad_and_normalize 保证任意尺寸输入都能安全进入窗口注意力;其训练入口尚未实现,这一点在使用 Swin2SRForImageSuperResolution 时需要明确预期。

登录后查看全文
热门项目推荐
相关项目推荐