Transformers 中的 ViT MSN:掩码孪生网络自监督预训练模型的全解析与图像分类实战
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
ViT MSN(Masked Siamese Networks,掩码孪生网络)是面向"标签高效学习"(label-efficient learning)的 Vision Transformer(ViT)自监督预训练方案,其核心思路是让模型把被随机遮挡 patch 的图像视图与未被遮挡的原始图像视图分配到的原型(prototype)对齐,从而学到高语义层级的图像表示。本文以 Transformers 官方模型文档为主体,结合仓库内 配置实现、模型实现、checkpoint 转换脚本 与测试用例,系统讲解 MSN 的原理、ViTMSNConfig全部配置项、ViTMSNModel/ViTMSNForImageClassification两大数据类用法、SDPA/FlashAttention 加速技巧以及下游微调与 checkpoint 转换的完整路径。
MSN 是什么:面向低标注量场景的自监督表示学习方法
ViTMSN 模型出自论文Masked Siamese Networks for Label-Efficient Learning(Assran、Caron、Misra 等人,2022)。从论文摘要可以看出其方法定位:
该方法把"包含随机掩码 patch 的图像视图"的表示,与"未掩码的原始图像"的表示进行匹配。这种自监督预训练策略在 Vision Transformer 上尤其具备可扩展性——因为网络中实际处理的只有未掩码的 patch,因此 MSN 提升了 joint-embedding 架构的可扩展性,同时能产出语义层级很高、在少样本(low-shot)图像分类上极具竞争力的表示。
代表性数据点是论文在 ImageNet-1K 上报告的结果:仅用5,000 张标注图像,MSN base 模型即可取得72.4% top-1 准确率;当标注量提高到ImageNet-1K 的 1%时,top-1 准确率可提升到75.7%,在当时刷新了该基准上自监督学习的最新记录。这一特性使 MSN 特别适用于低样本(low-shot)与极端低样本(extreme low-shot)两种数据稀缺场景。
Transformers 仓库在 2022-09-22 合入了这一模型(对应论文发布于 2022-04-14)。模型实现采用"模块化继承"方式:在 modular_vit_msn.py 中,ViTMSNPatchEmbeddings、ViTMSNAttention、ViTMSNMLP、ViTMSNLayer直接继承自vit.modeling_vit的对应类,然后通过 modeling_vit_msn.py 自动生成最终文件——因此整个编码器结构复用标准 ViT,差异点集中在 Embedding 层与初始化策略上。
使用要点:能直接用 backbone,也要知道局限
模型文档给出了三个核心使用提示,理解它们能避免踩坑:
- MSN 是一种自监督预训练方法:预训练目标是把"未掩码图像视图"分配到的原型与"同一图像掩码视图"的原型对齐。换言之,官方发布的是预训练好的特征提取骨干网络,而不是开箱即用的分类模型。
- 官方只发布了 ImageNet-1K 预训练的 backbone 权重:要在自己的图像分类数据集上使用,应当从
ViTMSNModel派生出ViTMSNForImageClassification,即在其上接一个分类头做微调。 - MSN 的甜区是低标注量场景:微调时仅使用 ImageNet-1K 1% 的标签即可达到 75.7% top-1 准确率。
此外需要注意一个架构细节:与常规 ViT 使用随机高斯(randn)初始化cls_token与position_embeddings不同,ViT MSN 对这两者以及可选的mask_token一律采用零初始化。这一点在 modular_vit_msn.py 的类注释与_init_weights中被显式标注,属于 MSN 与原始 ViT 的刻意差异,在加载官方权重时必须保持一致。
快速上手:加载 backbone 做特征提取
ViTMSNModel是骨干模型,输入pixel_values,输出last_hidden_state。modeling 文件中自带的示例即展示了完整调用链路(modeling_vit_msn.py 中ViTMSNModel.forward的 docstring 示例):
>>> from transformers import AutoImageProcessor, ViTMSNModel >>> import torch >>> from PIL import Image >>> import httpx >>> from io import BytesIO >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg" >>> with httpx.stream("GET", url) as response: ... image = Image.open(BytesIO(response.read())) >>> image_processor = AutoImageProcessor.from_pretrained("facebook/vit-msn-small") >>> model = ViTMSNModel.from_pretrained("facebook/vit-msn-small") >>> inputs = image_processor(images=image, return_tensors="pt") >>> with torch.no_grad(): ... outputs = model(**inputs) >>> last_hidden_states = outputs.last_hidden_state前向传播返回的是BaseModelOutput,其中last_hidden_state形状为(batch_size, num_patches + 1, hidden_size)——多出的 1 个 token 即[CLS]。若传入掩码位置张量bool_masked_pos,Embedding 层会把对应 patch 替换成可学习的mask_token(见 modeling_vit_msn.py 中ViTMSNEmbeddings.forward的掩码逻辑),这正是复现 MSN 预训练/微调流程时需要用到的底层能力。
图像分类:从 backbone 到下游分类头
作者未发布带分类头的权重,因此针对自己的分类数据集,应使用ViTMSNForImageClassification从ViTMSNModel初始化并微调。其底层结构(modeling_vit_msn.py 中ViTMSNForImageClassification)非常清晰:内部持有一个ViTMSNModel作为self.vit,分类头是nn.Linear(config.hidden_size, config.num_labels)(num_labels <= 0时为nn.Identity),前向时取序列第 0 个 token([CLS])的表示过分类头得到logits;传入labels时自动计算分类损失并随ImageClassifierOutput一并返回。
用官方权重直接做推理的示例:
>>> from transformers import AutoImageProcessor, ViTMSNForImageClassification >>> import torch >>> from PIL import Image >>> import httpx >>> from io import BytesIO >>> torch.manual_seed(2) >>> url = "http://images.cocodataset.org/val2017/000000039769.jpg" >>> with httpx.stream("GET", url) as response: ... image = Image.open(BytesIO(response.read())).convert("RGB") >>> image_processor = AutoImageProcessor.from_pretrained("facebook/vit-msn-small") >>> model = ViTMSNForImageClassification.from_pretrained("facebook/vit-msn-small") >>> inputs = image_processor(images=image, return_tensors="pt") >>> with torch.no_grad(): ... logits = model(**inputs).logits >>> # 模型预测 ImageNet 1000 类中的某一类 >>> predicted_label = logits.argmax(-1).item() >>> print(model.config.id2label[predicted_label]) tusker上述 cat 图片的tusker预测结果被固化在 model docstring 中,并且仓库的集成测试也对这一输出做了数值断言(见下文"测试验证"章节)。
在自有数据集上微调
官方推荐的微调路线是 examples/pytorch/image-classification 目录下的两个脚本:
run_image_classification.py:基于Trainer的标准训练脚本;run_image_classification_no_trainer.py:不依赖Trainer的轻量版本。
二者都通过--model_name_or_path/--dataset_name等参数驱动(参数说明见该目录下的 README.md),因此把 backbone 换成facebook/vit-msn-base之类的 MSN 权重即可复用完整训练链路。一个典型的数据集(Hub 数据集)训练命令形如:
python run_image_classification.py \ --model_name_or_path facebook/vit-msn-base \ --dataset_name beans \ --output_dir vit-msn-beans \ --remove_unused_columns false \ --do_train --do_eval \ --learning_rate 2e-4 --num_train_epochs 5 \ --per_device_train_batch_size 16 --per_device_eval_batch_size 16 \ --overwrite_output_dir --push_to_hub_model_id vit-msn-beans使用自有本地数据时,可参考同一 README 中的自定义数据集章节组织目录结构,通过--train_dir/--validation_dir传入图片目录。关于图像分类任务的通用预处理流程(图像处理器、数据增强等),可继续阅读 图像分类任务指南。
从源码理解关键机制:掩码、位置编码插值与双向注意力
对照 modeling_vit_msn.py 可以梳理出四个值得理解实现细节:
- Patch Embedding(
ViTMSNPatchEmbeddings):用nn.Conv2d(config.num_channels, config.hidden_size, kernel_size=patch_size, stride=patch_size)把(batch, 3, H, W)的像素图切成(H/patch_size) × (W/patch_size)个 patch,num_patches即 patch 总数;forward 中会校验输入通道数是否等于配置的num_channels。 - 掩码 token 机制(
ViTMSNEmbeddings.forward):当传入bool_masked_pos(形状(batch_size, num_patches),1 表示掩码、0 表示保留)时,被掩码 patch 的嵌入被mask_token替换——这正是 MSN 在微调/评估阶段模拟"掩码视图"的入口。注意该能力只有use_mask_token=True构造的ViTMSNModel才具备(默认False,此时mask_token为None)。 - 位置编码插值(
interpolate_pos_encoding):当推理图像分辨率与训练分辨率不一致时,interpolate_pos_encoding=True会用 bicubic 插值把预训练位置编码重采样到(H/patch_size, W/patch_size)的网格上,从而支持更高分辨率输入;若关闭插值而输入尺寸又对不上,会抛出尺寸不匹配的ValueError。该方法同时兼容torch.jit跟踪导出。 - 双向注意力与注意力后端:ViT MSN 是编码器结构、非因果,attention mask 通过
create_bidirectional_mask生成。注意力实现走统一的ALL_ATTENTION_FUNCTIONS接口(ViTMSNAttention.forward中get_interface(self.config._attn_implementation, eager_attention_fal_forward)),与仓库通用的 FlashAttention / SDPA / FlexAttention 后端子模块打通,模型类声明了_supports_sdpa = True、_supports_flash_attn = True、_supports_flex_attn = True以及supports_gradient_checkpointing = True,说明 MSN 直接继承了 ViT 家族的全部注意力加速与显存优化能力。
ViTMSNConfig 配置项详解
ViTMSNConfig(configuration_vit_msn.py)的model_type为"vit_msn",结构上等同于 ViT 配置。下表汇总了各字段及其仓库中的默认值:
| 配置字段 | 默认值 | 含义 |
|---|---|---|
hidden_size | 768 | 隐藏层维度(base 规模) |
num_hidden_layers | 12 | Transformer 编码器层数 |
num_attention_heads | 12 | 注意力头数 |
intermediate_size | 3072 | MLP 中间层维度 |
hidden_act | "gelu" | 隐藏层激活函数 |
hidden_dropout_prob | 0.0 | 隐藏层 Dropout 概率 |
attention_probs_dropout_prob | 0.0 | 注意力概率 Dropout |
initializer_range | 0.02 | 权重初始化标准差范围 |
layer_norm_eps | 1e-6 | LayerNorm epsilon |
image_size | 224 | 输入图像尺寸,可为int或(H, W) |
patch_size | 16 | patch 尺寸,可为int或(H, W) |
num_channels | 3 | 输入图像通道数 |
qkv_bias | True | Q/K/V 线性投影是否带偏置 |
s/16、b/16、l/16等不同规模 checkpoint 对应的差异正是在此配置上体现:例如 small 为hidden_size=384、intermediate_size=1536、6 头;large 为hidden_size=1024、intermediate_size=4096、24 层、16 头并把hidden_dropout_prob调为0.1(这些取值可直接在 convert_msn_to_pytorch.py 的convert_vit_msn_checkpoint中看到)。此外分类所需的num_labels、id2label、label2id等属性继承自PreTrainedConfig,由对应 checkpoint 或下游任务自行设置(转换脚本中转换 ImageNet 分类权重时即显式config.num_labels = 1000并加载imagenet-1k-id2label.json)。
用 SDPA 加速推理:实测数据与开关方式
自 PyTorch 2.1.1 起,当对应实现可用时,模型默认启用PyTorch 原生缩放点积注意力(SDPA,torch.nn.functional.scaled_dot_product_attention);也可以通过在from_pretrained()中显式传attn_implementation="sdpa"强制使用:
from transformers import ViTMSNForImageClassification model = ViTMSNForImageClassification.from_pretrained( "facebook/vit-msn-base", attn_implementation="sdpa", device_map="auto" )为获得最佳加速效果,官方建议把模型加载为半精度(torch.float16或torch.bfloat16)。模型文档给出了一组本地基准数据(A100-40GB、PyTorch 2.3.0、Ubuntu 22.04、float32、facebook/vit-msn-base推理):
| Batch size | eager 平均推理时间 (ms) | sdpa 平均推理时间 (ms) | 加速比 (Sdpa / Eager, x) |
|---|---|---|---|
| 1 | 7 | 6 | 1.17 |
| 2 | 8 | 6 | 1.33 |
| 4 | 8 | 6 | 1.33 |
| 8 | 8 | 6 | 1.33 |
需要说明的是,该表格是模型文档在特定软硬件组合下的实测参考值;实际加速幅度取决于 GPU 型号、PyTorch 版本、批大小与精度,应以上述方式在自己的环境复测为准。除sdpa外,由于_supports_flash_attn = True,同样可在支持的硬件上通过attn_implementation="flash_attention_2"启用 FlashAttention。
把官方 MSN 权重导入 Transformers:转换脚本机制
ViTMSNPreTrainedModel的base_model_prefix为"vit",意味着官方 Facebook 权重(key 形如module.blocks.*)需要经过键名重映射才能被 Transformers 正常加载。convert_msn_to_pytorch.py 正是这一"桥接"工具,它的处理逻辑本身也揭示了原版实现与 Transformers 实现之间的映射关系:
- 键名重命名(
create_rename_keys):把module.blocks.{i}.norm1/attn.proj/norm2/mlp.fc1/mlp.fc2等映射为vit.encoder.layer.{i}.layernorm_before / attention.output.dense / layernorm_after / intermediate.dense / output.dense等 Transformers 命名; - QKV 矩阵切分(
read_in_q_k_v):原版 timm 风格把 query/key/value 打包在单个attn.qkv矩阵中,脚本按hidden_size边界切分为独立的 Q、K、V 权重与偏置; - 结构适配:加载到
ViTMSNModel时剥离预训练专用的三阶段投影头module.fc.*(含 BatchNorm 层,remove_projection_head),加载分类头时则只保留norm与head; - 结果校验:转换后会用 COCO 样例图跑一次前向,并把
last_hidden_state的起始切片与各规模 checkpoint 的参考值做allclose(atol=1e-4)比对,确保转换无损。
脚本同时会根据 checkpoint URL 中的s16/l16/b4/l7等标识自动调整配置(hidden_size、patch_size、层数等),支持多种规模权重的一次性导入。
测试验证:数值断言与形状契约
tests/models/vit_msn/test_modeling_vit_msn.py 对模型契约做了完整约束,可作为使用时的行为参照:
- 输出形状:
ViTMSNModel输出last_hidden_state形状必须为(batch_size, num_patches + 1, hidden_size),ViTMSNForImageClassification输出 logits 形状为(batch_size, num_labels),其中num_patches = (image_size / patch_size) ^ 2; - 灰度图支持:测试显式验证了把
num_channels设为 1 后模型仍能正常出 logits; - 无文本模态:由于 MSN 不使用
input_ids/inputs_embeds,相应通用测试被跳过(has_text_modality=False); - pipeline 映射:
ViTMSNModel支持image-feature-extractionpipeline,ViTMSNForImageClassification支持image-classificationpipeline; - 集成测试数值断言:slow 测试加载
facebook/vit-msn-small在 COCO 示例图上推理,断言 logits 形状为(1, 1000),并校验 logits 起始三个元素与期望值[0.5588, 0.6853, -0.5929]在rtol=1e-4, atol=1e-4内吻合——这与上文 docstring 中tusker的预测示例是同一组权重与图片的互相印证。
更多学习资源
- 模型本体文档:ViTMSN 官方文档页;
- 骨干与分类头实现:modeling_vit_msn.py,模块化源文件为 modular_vit_msn.py;
- 配置定义:configuration_vit_msn.py;
- checkpoint 转换工具:convert_msn_to_pytorch.py;
- 微调脚本与完整参数说明:examples/pytorch/image-classification 下的
run_image_classification.py与 README.md; - 图像分类任务通用指南:docs/source/en/tasks/image_classification.md。
总体而言,在 Transformers 生态中使用 ViT MSN 的路径非常清晰:facebook/vit-msn-{small,base,large}等 Hub 权重提供了高质量自监督骨干,ViTMSNModel承担特征提取,ViTMSNForImageClassification通过极少量标注即可微调出在低标注量场景下有竞争力的分类模型,而 SDPA/FlashAttention 与半精度加载则让它在现代 GPU 上的推理开销保持在可控范围。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考