尧图精选

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

🕒 发布时间:2026/9/9 23:32:23 📁 来源:尧图网络
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/transformersViT MSNMasked Siamese Networks掩码孪生网络是面向标签高效学习label-efficient learning的 Vision TransformerViT自监督预训练方案其核心思路是让模型把被随机遮挡 patch 的图像视图与未被遮挡的原始图像视图分配到的原型prototype对齐从而学到高语义层级的图像表示。本文以 Transformers 官方模型文档为主体结合仓库内 配置实现、模型实现、checkpoint 转换脚本 与测试用例系统讲解 MSN 的原理、ViTMSNConfig全部配置项、ViTMSNModel/ViTMSNForImageClassification两大数据类用法、SDPA/FlashAttention 加速技巧以及下游微调与 checkpoint 转换的完整路径。MSN 是什么面向低标注量场景的自监督表示学习方法ViTMSN 模型出自论文Masked Siamese Networks for Label-Efficient LearningAssran、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(imagesimage, return_tensorspt) 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_posEmbedding 层会把对应 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(imagesimage, return_tensorspt) 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 EmbeddingViTMSNPatchEmbeddings用nn.Conv2d(config.num_channels, config.hidden_size, kernel_sizepatch_size, stridepatch_size)把(batch, 3, H, W)的像素图切成(H/patch_size) × (W/patch_size)个 patchnum_patches即 patch 总数forward 中会校验输入通道数是否等于配置的num_channels。掩码 token 机制ViTMSNEmbeddings.forward当传入bool_masked_pos形状(batch_size, num_patches)1 表示掩码、0 表示保留时被掩码 patch 的嵌入被mask_token替换——这正是 MSN 在微调/评估阶段模拟掩码视图的入口。注意该能力只有use_mask_tokenTrue构造的ViTMSNModel才具备默认False此时mask_token为None。位置编码插值interpolate_pos_encoding当推理图像分辨率与训练分辨率不一致时interpolate_pos_encodingTrue会用 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 配置项详解ViTMSNConfigconfiguration_vit_msn.py的model_type为vit_msn结构上等同于 ViT 配置。下表汇总了各字段及其仓库中的默认值配置字段默认值含义hidden_size768隐藏层维度base 规模num_hidden_layers12Transformer 编码器层数num_attention_heads12注意力头数intermediate_size3072MLP 中间层维度hidden_actgelu隐藏层激活函数hidden_dropout_prob0.0隐藏层 Dropout 概率attention_probs_dropout_prob0.0注意力概率 Dropoutinitializer_range0.02权重初始化标准差范围layer_norm_eps1e-6LayerNorm epsilonimage_size224输入图像尺寸可为int或(H, W)patch_size16patch 尺寸可为int或(H, W)num_channels3输入图像通道数qkv_biasTrueQ/K/V 线性投影是否带偏置s/16、b/16、l/16等不同规模 checkpoint 对应的差异正是在此配置上体现例如 small 为hidden_size384、intermediate_size1536、6 头large 为hidden_size1024、intermediate_size4096、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 原生缩放点积注意力SDPAtorch.nn.functional.scaled_dot_product_attention也可以通过在from_pretrained()中显式传attn_implementationsdpa强制使用from transformers import ViTMSNForImageClassification model ViTMSNForImageClassification.from_pretrained( facebook/vit-msn-base, attn_implementationsdpa, device_mapauto )为获得最佳加速效果官方建议把模型加载为半精度torch.float16或torch.bfloat16。模型文档给出了一组本地基准数据A100-40GB、PyTorch 2.3.0、Ubuntu 22.04、float32、facebook/vit-msn-base推理Batch sizeeager 平均推理时间 (ms)sdpa 平均推理时间 (ms)加速比 (Sdpa / Eager, x)1761.172861.334861.338861.33需要说明的是该表格是模型文档在特定软硬件组合下的实测参考值实际加速幅度取决于 GPU 型号、PyTorch 版本、批大小与精度应以上述方式在自己的环境复测为准。除sdpa外由于_supports_flash_attn True同样可在支持的硬件上通过attn_implementationflash_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 的参考值做allcloseatol1e-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_modalityFalsepipeline 映射ViTMSNModel支持image-feature-extractionpipelineViTMSNForImageClassification支持image-classificationpipeline集成测试数值断言slow 测试加载facebook/vit-msn-small在 COCO 示例图上推理断言 logits 形状为(1, 1000)并校验 logits 起始三个元素与期望值[0.5588, 0.6853, -0.5929]在rtol1e-4, atol1e-4内吻合——这与上文 docstring 中tusker的预测示例是同一组权重与图片的互相印证。更多学习资源模型本体文档ViTMSN 官方文档页骨干与分类头实现modeling_vit_msn.py模块化源文件为 modular_vit_msn.py配置定义configuration_vit_msn.pycheckpoint 转换工具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),仅供参考
上一篇/下一篇内容由系统自动关联 返回资讯列表 →