news 2026/9/9 23:32:19

Transformers 中的 ViT MSN:掩码孪生网络自监督预训练模型的全解析与图像分类实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Transformers 中的 ViT MSN:掩码孪生网络自监督预训练模型的全解析与图像分类实战

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 中,ViTMSNPatchEmbeddingsViTMSNAttentionViTMSNMLPViTMSNLayer直接继承自vit.modeling_vit的对应类,然后通过 modeling_vit_msn.py 自动生成最终文件——因此整个编码器结构复用标准 ViT,差异点集中在 Embedding 层与初始化策略上。

使用要点:能直接用 backbone,也要知道局限

模型文档给出了三个核心使用提示,理解它们能避免踩坑:

  1. MSN 是一种自监督预训练方法:预训练目标是把"未掩码图像视图"分配到的原型与"同一图像掩码视图"的原型对齐。换言之,官方发布的是预训练好的特征提取骨干网络,而不是开箱即用的分类模型。
  2. 官方只发布了 ImageNet-1K 预训练的 backbone 权重:要在自己的图像分类数据集上使用,应当从ViTMSNModel派生出ViTMSNForImageClassification,即在其上接一个分类头做微调。
  3. MSN 的甜区是低标注量场景:微调时仅使用 ImageNet-1K 1% 的标签即可达到 75.7% top-1 准确率。

此外需要注意一个架构细节:与常规 ViT 使用随机高斯(randn)初始化cls_tokenposition_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 到下游分类头

作者未发布带分类头的权重,因此针对自己的分类数据集,应使用ViTMSNForImageClassificationViTMSNModel初始化并微调。其底层结构(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 可以梳理出四个值得理解实现细节:

  1. Patch EmbeddingViTMSNPatchEmbeddings):用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
  2. 掩码 token 机制ViTMSNEmbeddings.forward):当传入bool_masked_pos(形状(batch_size, num_patches),1 表示掩码、0 表示保留)时,被掩码 patch 的嵌入被mask_token替换——这正是 MSN 在微调/评估阶段模拟"掩码视图"的入口。注意该能力只有use_mask_token=True构造的ViTMSNModel才具备(默认False,此时mask_tokenNone)。
  3. 位置编码插值interpolate_pos_encoding):当推理图像分辨率与训练分辨率不一致时,interpolate_pos_encoding=True会用 bicubic 插值把预训练位置编码重采样到(H/patch_size, W/patch_size)的网格上,从而支持更高分辨率输入;若关闭插值而输入尺寸又对不上,会抛出尺寸不匹配的ValueError。该方法同时兼容torch.jit跟踪导出。
  4. 双向注意力与注意力后端:ViT MSN 是编码器结构、非因果,attention mask 通过create_bidirectional_mask生成。注意力实现走统一的ALL_ATTENTION_FUNCTIONS接口(ViTMSNAttention.forwardget_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_size768隐藏层维度(base 规模)
num_hidden_layers12Transformer 编码器层数
num_attention_heads12注意力头数
intermediate_size3072MLP 中间层维度
hidden_act"gelu"隐藏层激活函数
hidden_dropout_prob0.0隐藏层 Dropout 概率
attention_probs_dropout_prob0.0注意力概率 Dropout
initializer_range0.02权重初始化标准差范围
layer_norm_eps1e-6LayerNorm epsilon
image_size224输入图像尺寸,可为int(H, W)
patch_size16patch 尺寸,可为int(H, W)
num_channels3输入图像通道数
qkv_biasTrueQ/K/V 线性投影是否带偏置

s/16b/16l/16等不同规模 checkpoint 对应的差异正是在此配置上体现:例如 small 为hidden_size=384intermediate_size=1536、6 头;large 为hidden_size=1024intermediate_size=4096、24 层、16 头并把hidden_dropout_prob调为0.1(这些取值可直接在 convert_msn_to_pytorch.py 的convert_vit_msn_checkpoint中看到)。此外分类所需的num_labelsid2labellabel2id等属性继承自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.float16torch.bfloat16)。模型文档给出了一组本地基准数据(A100-40GB、PyTorch 2.3.0、Ubuntu 22.04、float32facebook/vit-msn-base推理):

Batch sizeeager 平均推理时间 (ms)sdpa 平均推理时间 (ms)加速比 (Sdpa / Eager, x)
1761.17
2861.33
4861.33
8861.33

需要说明的是,该表格是模型文档在特定软硬件组合下的实测参考值;实际加速幅度取决于 GPU 型号、PyTorch 版本、批大小与精度,应以上述方式在自己的环境复测为准。除sdpa外,由于_supports_flash_attn = True,同样可在支持的硬件上通过attn_implementation="flash_attention_2"启用 FlashAttention。

把官方 MSN 权重导入 Transformers:转换脚本机制

ViTMSNPreTrainedModelbase_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),加载分类头时则只保留normhead
  • 结果校验:转换后会用 COCO 样例图跑一次前向,并把last_hidden_state的起始切片与各规模 checkpoint 的参考值做allcloseatol=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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/9 23:31:03

C#深度学习落地实践:ONNX Runtime+YOLOv8推理完整指南

简介&#xff1a;这是一份基于Visual Studio 2013开发的C#深度学习源码示例&#xff0c;面向希望在Windows环境中快速上手深度学习的C#工程师与学生。相比常见的Linux移植版本&#xff0c;它省去配置第三方库的难题&#xff0c;安装VS2013即可直接编译运行&#xff0c;大幅降低…

作者头像 李华
网站建设 2026/9/9 23:26:48

FPGA实战:BT656接口720x576格式的Verilog实现与时序仿真

简介&#xff1a;一份基于Verilog HDL的BT656视频编码实现&#xff0c;面向FPGA开发者和数字视频接口学习者&#xff0c;解决RGB888像素格式到BT656标准数据流的转换&#xff0c;并适配720x576分辨率输出。压缩包共130个文件&#xff0c;大小约4.14MB&#xff0c;核心包含bt656…

作者头像 李华
网站建设 2026/9/9 23:26:45

Changes Made

Changes Made 【免费下载链接】oh-my-claudecode Teams-first Multi-agent orchestration for Claude Code 项目地址: https://gitcode.com/GitHub_Trending/oh/oh-my-claudecode file.ts:42-55: [what changed and why] Verification Build: [command] -> [pass/f…

作者头像 李华
网站建设 2026/9/9 23:25:00

基于YOLOv8的AI蒸汽除草机器人:从Ubuntu环境到目标检测实战

各位关注 AI 与机器人方向的朋友们&#xff0c;大家好。今天我想和大家分享一个非常有“落地感”的 AI 项目&#xff1a;AI 蒸汽除草机器人。最近看到明尼苏达州发明家打造无化学除草机器人的相关消息&#xff0c;确实让人眼前一亮。在环保要求越来越高的背景下&#xff0c;用高…

作者头像 李华