Transformers 中的 BEiT 模型:BERT 式掩码图像建模预训练视觉 Transformer 的架构、配置与实战指南
【免费下载链接】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
BEiT(Bidirectional Encoder representation from Image Transformers)是首个让视觉 Transformer(ViT)的自监督预训练效果全面超越有监督预训练的代表性工作。本文以 Hugging Face Transformers 仓库中的 BEiT 模型文档 为核心骨架,结合 modeling_beit.py、configuration_beit.py 等源码实现,系统讲解 BEiT 的掩码图像建模原理、完整配置参数、图像预处理、四类下游任务模型以及 SDPA 注意力加速,帮助你直接上手调用预训练权重完成图像分类、语义分割与掩码图像建模等任务。
BEiT 模型概览:把 BERT 的自监督范式搬到视觉领域
BEiT 由 Hangbo Bao、Li Dong 和 Furu Wei 在论文BEiT: BERT Pre-Training of Image Transformers中提出,论文于 2021 年 6 月 15 日发布在 Hugging Face Papers 上,并于 2021 年 8 月 4 日合入本仓库。它的核心贡献在于:受 NLP 领域 BERT 的启发,BEiT 提出了**掩码图像建模(Masked Image Modeling, MIM)**任务来预训练视觉 Transformer,这也是第一个让 ViT 的自监督预训练在效果上超过有监督预训练的方法。
与 原始 ViT 论文 直接预测图像类别标签不同,BEiT 的预训练目标是根据被掩码的图像块(patch),预测来自 OpenAI DALL-E 模型 码本(codebook)的视觉 token。论文摘要中的关键表述如下:
我们引入了自监督视觉表示模型 BEiT。遵循 NLP 领域 BERT 的思路,我们提出掩码图像建模任务来预训练视觉 Transformer。具体地,每张图像在预训练中有两个视图:图像块(如 16×16 像素)与视觉 token(即离散 token)。我们先将原始图像"token 化"为视觉 token,然后随机掩码部分图像块并送入骨干 Transformer,预训练目标是基于被破坏的图像块恢复原始视觉 token。预训练完成后,通过在预训练编码器上追加任务层,直接在下游任务上微调模型参数。图像分类与语义分割的实验表明,我们的模型取得了具有竞争力的结果:例如 base 尺寸的 BEiT 在 ImageNet-1K 上达到 83.2% 的 top-1 准确率,明显超过相同设置下从零训练的 DeiT(81.8%);large 尺寸的 BEiT 仅使用 ImageNet-1K 就达到 86.3%,甚至超过了在 ImageNet-22K 上有监督预训练的 ViT-L(85.2%)。
两视图预训练:图像块 + 视觉 token
从论文摘要和仓库源码都可以还原出 BEiT 预训练的完整流程,它包含两条并行信息流:
- 图像块视图:将原图切成 16×16 像素的 patch(base 配置下 224×224 输入会产生 196 个 patch),由 BeitPatchEmbeddings 用
nn.Conv2d(num_channels, hidden_size, kernel_size=patch_size, stride=patch_size)一次性完成投影,输出形状为(batch_size, seq_length, hidden_size)的序列特征; - 视觉 token 视图:借助 DALL-E 的离散图像 tokenizer(VQ-VAE),把图像编码为离散的视觉 token,作为预训练阶段需要"恢复"的监督信号。
训练时随机掩码一部分 patch,将其替换为可学习的 mask token(见 BeitEmbeddings.forward 中bool_masked_pos与mask_token的替换逻辑),把被破坏的 patch 序列喂给骨干 Transformer,最后用语言建模头在掩码位置预测对应的视觉 token。注意,由于预测的是 VQ-VAE 视觉 token 而非 RGB 像素值,BEiT 的掩码图像建模实现与AutoModelForMaskedImageModeling不兼容,需要使用BeitForMaskedImageModeling直接调用(这一点在 modeling_beit.py 的类文档中有明确说明)。
架构与源码实现:基于 ViT 骨架的模块化组装
BEiT 在架构上就是一个常规的视觉 Transformer,区别仅在于预训练方式。仓库采用模块化(modular)方式组织代码:modular_beit.py 是源文件,它直接复用了ViTPatchEmbeddings、ViTEmbeddings、ViTAttention、ViTMLP、ViTLayer、ViTPreTrainedModel等 ViT 组件,仅替换了BeitRelativePositionBias等 BEiT 特有模块;modeling_beit.py 则是自动生成的实际模型实现文件(文件头部有"由 modular 文件生成,请勿手动编辑"的警告)。
前向流程与关键模块
BeitModel(modeling_beit.py)的完整前向链路如下:
BeitEmbeddings:patch 投影 → (可选)掩码替换 → 拼接[CLS]token → (可选)添加绝对位置编码 → dropout;- 每个
BeitLayer(对应 timm 实现的 Block):Pre-LayerNorm → 多头自注意力 → layer scale(lambda_1)→ DropPath + 残差 → Post-LayerNorm → MLP → layer scale(lambda_2)→ DropPath + 残差; - 末端
layernorm与BeitPooler汇聚出序列输出与池化输出。
值得一提的实现细节:
- Layer scale 初始化:
BeitLayer中的lambda_1/lambda_2以config.layer_scale_init_value(默认 0.1)初始化(modeling_beit.py); - DropPath(随机深度):
BeitDropPath按层线性递增 dropout 率(drop_path_rate * i / (num_hidden_layers - 1)); - CLS 与均值池化:BeitPooler 在
use_mean_pooling=True时对 patch token(剔除[CLS])做 LayerNorm 后的均值池化,否则取[CLS]token 的最终隐藏状态; - 输出结构:
BeitModelOutputWithPooling(modeling_beit.py)在标准BaseModelOutputWithPooling基础上明确了pooler_output的语义(均值池化或 CLS),并支持hidden_states/attentions的逐层输出。
T5 式相对位置偏置
BEiT 使用受 T5 模型启发的相对位置嵌入:预训练时作者在多个自注意力层之间共享相对位置偏置(use_shared_relative_position_bias),微调时每一层的相对位置偏置用预训练得到的共享偏置初始化。BeitRelativePositionBias 实现了完整的相对位置索引生成与偏置表插值逻辑,支持任意窗口尺寸(输入分辨率变化时通过双线性插值扩展偏置表)。源码中_keys_to_ignore_on_load_unexpected = [r".*relative_position_index.*"](modeling_beit.py)表明相对位置索引是运行时动态生成的缓存张量,不属于需要加载的模型权重。
使用要点:如果想从零预训练BEiT,必须将BeitConfig的use_relative_position_bias(每层独立偏置)或use_shared_relative_position_bias(跨层共享偏置)设置为True,否则模型中不会包含位置嵌入。日常加载官方预训练权重做微调时,权重自带相对位置偏置,无需额外配置。
BeitConfig:完整配置参数详解
BeitConfig定义在 configuration_beit.py 中,继承自PreTrainedConfig与BackboneConfigMixin。下表整理了当前仓库中全部配置项及其默认值:
| 配置项 | 默认值 | 说明 |
|---|---|---|
vocab_size | 8192 | 视觉 token 码本大小(即掩码图像建模头的输出维度) |
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-12 | LayerNorm epsilon |
image_size | 224 | 输入图像分辨率(可传int或(h, w)元组) |
patch_size | 16 | patch 尺寸 |
num_channels | 3 | 输入通道数 |
use_mask_token | False | 是否为掩码图像建模启用可学习 mask token |
use_absolute_position_embeddings | False | 是否使用绝对位置编码 |
use_relative_position_bias | False | 是否在注意力层中使用 T5 式相对位置偏置 |
use_shared_relative_position_bias | False | 是否跨层共享同一份相对位置偏置 |
layer_scale_init_value | 0.1 | layer scale 初始化值(≤0 时禁用) |
drop_path_rate | 0.1 | 随机深度最大丢弃率 |
use_mean_pooling | True | 分类头之前对 patch token 均值池化(否则用 CLS) |
pool_scales | (1, 2, 3, 6) | 语义分割 PSP 模块的池化尺度 |
use_auxiliary_head | True | 训练时是否使用辅助分割头 |
auxiliary_loss_weight | 0.4 | 辅助头交叉熵损失的权重 |
auxiliary_channels | 256 | 辅助头通道数 |
auxiliary_num_convs | 1 | 辅助头卷积层数 |
auxiliary_concat_input | False | 分类层前是否拼接辅助头输入 |
semantic_loss_ignore_index | 255 | 语义分割损失的忽略索引 |
add_fpn | False | 作为骨干网时是否附加 FPN(仅BeitBackbone使用) |
reshape_hidden_states | True | 骨干输出是否重排为 4D 特征图 |
从源码可以看到两个联动校验规则(configuration_beit.py):add_fpn=True时out_indices必须恰好为 4 个整数(base 架构建议[3, 5, 7, 11]);out_indices也可用旧参数名segmentation_indices传入并自动转换。stage_names由["stem"] + ["stage1"..."stage12"]组成,供骨干输出对齐使用。
初始化一个 BEiT 配置与随机权重模型的标准写法:
>>> from transformers import BeitConfig, BeitModel >>> # 初始化 beit-base-patch16-224-pt22k 风格的配置 >>> configuration = BeitConfig() >>> # 从配置初始化(随机权重)模型 >>> model = BeitModel(configuration) >>> # 访问模型配置 >>> configuration = model.config图像预处理:BeitImageProcessor 与 BeitImageProcessorPil
由于 BEiT 模型要求每张输入图像具有相同分辨率,必须使用图像处理器完成 resize(或 rescale)与归一化。仓库提供两个后端实现:
- BeitImageProcessor:基于 Torchvision 后端的处理器;
- BeitImageProcessorPil:基于 PIL 的处理器,二者共享
BeitImageProcessorKwargs。
它们的默认预处理参数完全一致(image_processing_beit.py):
| 参数 | 默认值 |
|---|---|
resample | BICUBIC |
image_mean/image_std | ImageNet 标准均值/标准差 |
size | 224×224 |
crop_size | 224×224 |
do_resize | True |
do_center_crop | False |
do_rescale | True |
do_normalize | True |
do_reduce_labels | False |
preprocess方法除了处理普通图像外,还支持传入segmentation_maps:分割标签图会以do_normalize=False, do_rescale=False独立处理,避免归一化破坏类别标签的整数语义,并转为int64张量(image_processing_beit.py)。do_reduce_labels用于 ADE20k 这类以 0 表示背景、但背景不参与类别计数的数据集——开启后所有标签值减 1,背景被替换为 255(忽略索引)。post_process_semantic_segmentation则把BeitForSemanticSegmentation的原始 logits 转换为逐像素的分割图(支持target_sizes尺寸还原与return_segmentation_scores概率输出)。
下游任务:四类模型 + 骨干网络
BeitForImageClassification:图像分类
BeitForImageClassification 在BeitModel(带池化层)之上接一个线性分类头。分类头输入是 patch token 均值池化(或 CLS)后的pooler_output,num_labels == 1时计算 MSE 回归损失,否则计算交叉熵损失。推断示例:
>>> from transformers import AutoImageProcessor, BeitForImageClassification >>> from PIL import Image >>> image_processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-patch16-224") >>> model = BeitForImageClassification.from_pretrained("microsoft/beit-base-patch16-224") >>> inputs = image_processor(images=image, return_tensors="pt") >>> outputs = model(**inputs) >>> logits = outputs.logitsBeitForMaskedImageModeling:掩码图像建模
BeitForMaskedImageModeling 在骨干之上叠加 LayerNorm 与lm_head(nn.Linear(hidden_size, vocab_size)),预测被掩码 patch 的视觉 token。它通过bool_masked_pos(形状(batch_size, num_patches),1 表示被掩码)指定掩码位置,仅在掩码位置计算交叉熵损失:
>>> from transformers import AutoImageProcessor, BeitForMaskedImageModeling >>> import torch >>> from PIL import Image >>> image_processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-patch16-224-pt22k") >>> model = BeitForMaskedImageModeling.from_pretrained("microsoft/beit-base-patch16-224-pt22k") >>> num_patches = (model.config.image_size // model.config.patch_size) ** 2 >>> pixel_values = image_processor(images=image, return_tensors="pt").pixel_values >>> # 生成 (1, num_patches) 的随机布尔掩码 >>> bool_masked_pos = torch.randint(low=0, high=2, size=(1, num_patches)).bool() >>> outputs = model(pixel_values, bool_masked_pos=bool_masked_pos) >>> loss, logits = outputs.loss, outputs.logits >>> list(logits.shape) [1, 196, 8192]输出 logits 形状[1, 196, 8192]对应 196 个 patch 与 8192 的码本大小,与vocab_size配置严格对应。预训练或继续预训练时用labels(视觉 token ID)监督掩码位置即可。
BeitForSemanticSegmentation:语义分割
BeitForSemanticSegmentation 是 BEiT 在密集预测任务上的完整实现,结构上由三部分构成:
- FPN 颈部(BeitFPNNeck):把
out_indices选出的 4 层特征图映射为 4 级金字塔(2 倍上采样 / 4 倍上采样 / 恒等 / 2 倍下采样,见 modeling_beit.py); - 解码头(BeitUperHead):UPerNet 风格的 PSP 模块(
pool_scales金字塔池化)+ FPN 融合,最后用 1×1 卷积输出num_labels通道(modeling_beit.py); - 辅助头(BeitFCNHead):FCN 风格的辅助监督,其通道数、卷积层数、权重等由
auxiliary_channels/auxiliary_num_convs/auxiliary_loss_weight控制。
该模型强制要求config.out_indices为恰好 4 个整数(base 架构用[3, 5, 7, 11]),否则直接抛出ValueError。推断时 logits 形状为(batch_size, num_labels, height, width):
>>> from transformers import AutoImageProcessor, BeitForSemanticSegmentation >>> from PIL import Image >>> image_processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-finetuned-ade-640-640") >>> model = BeitForSemanticSegmentation.from_pretrained("microsoft/beit-base-finetuned-ade-640-640") >>> inputs = image_processor(images=image, return_tensors="pt") >>> outputs = model(**inputs) >>> logits = outputs.logits # (batch_size, num_labels, height, width)BeitBackbone:供 DETR / MaskFormer 等框架使用
BeitBackbone 将 BEiT 作为通用骨干网络暴露给 DETR、MaskFormer 等检测/分割框架,支持通过out_features/out_indices选择输出层,reshape_hidden_states控制输出 4D 特征图或 3D 序列,add_fpn=True时附加 FPN。可结合AutoBackbone使用:
>>> from transformers import AutoImageProcessor, AutoBackbone >>> processor = AutoImageProcessor.from_pretrained("microsoft/beit-base-patch16-224") >>> model = AutoBackbone.from_pretrained( ... "microsoft/beit-base-patch16-224", out_features=["stage1", "stage2", "stage3", "stage4"] ... ) >>> outputs = model(**processor(image, return_tensors="pt")) >>> list(outputs.feature_maps[-1].shape) [1, 768, 14, 14]检查点命名与可用权重
理解检查点命名规则有助于正确选择预训练权重。BEiT 每个检查点的名字都同时反映了预训练/微调时使用的 patch 分辨率与图像分辨率:
microsoft/beit-base-patch16-224:base 尺寸架构,patch 分辨率 16×16,微调分辨率 224×224;microsoft/beit-base-patch16-224-pt22k:在 ImageNet-22k 上预训练(仅预训练);microsoft/beit-large-patch16-224-pt22k-ft22k:在 ImageNet-22k 预训练并在 ImageNet-22k 上微调;microsoft/beit-base-finetuned-ade-640-640:在 ADE20k 上微调至 640×640 的语义分割权重。
可用检查点分三类:① 仅在 ImageNet-22k(约 1400 万张图像、2.2 万类)上预训练;② 在 ImageNet-22k 上进一步微调;③ 在 ImageNet-1k(ILSVRC 2012,约 130 万张图像、1000 类)上微调。仓库测试 test_modeling_beit.py 的 slow 用例覆盖了上述全部检查点类型,可作为挑选权重的参考依据。
注意力加速:Scaled Dot Product Attention(SDPA)
当前仓库的 BEiT 实现已全面接入 PyTorch 原生 SDPA。源码层面,BeitPreTrainedModel声明了_supports_sdpa = True与_supports_flash_attn = False(modeling_beit.py),BeitAttention通过ALL_ATTENTION_FUNCTIONS.get_interface(config._attn_implementation, eager_attention_forward)按配置分发到 eager 或 SDPA 实现(modeling_beit.py)。
当 PyTorch 版本 ≥ 2.1.1 且硬件可用时,SDPA 默认启用;也可以显式指定attn_implementation="sdpa":
from transformers import BeitForImageClassification model = BeitForImageClassification.from_pretrained( "microsoft/beit-base-patch16-224", attn_implementation="sdpa", device_map="auto" )官方英文文档记录了在 NVIDIA GeForce RTX 2060-8GB、PyTorch 2.5.1、Ubuntu 20.04 环境下、float16精度 +microsoft/beit-base-patch16-224的本地基准结果(详见 英文版 BEiT 文档):
训练场景(50 步、batch=2、图像 1048×640):每 batch 耗时从 eager 的 0.984s 降至 SDPA 的 0.746s,加速约 31.98%;峰值显存从 6738.9MB 降至 4319.9MB,节省约 56%。
推理场景(各 batch 下的对比):
| 图像 batch 数 | Eager (s/iter) | SDPA (s/iter) | 加速比 | 显存节省 |
|---|---|---|---|---|
| 1 | 0.012 | 0.011 | 1.05× | 0.24% |
| 4 | 0.013 | 0.011 | 1.18× | 3.23% |
| 16 | 0.045 | 0.035 | 1.30× | 10.08% |
| 32 | 0.088 | 0.066 | 1.33× | 17.04% |
可见 batch 越大、分辨率越高,SDPA 的收益越明显。文档同时建议:为获得最佳加速效果,请以半精度加载模型(torch.float16或torch.bfloat16)。测试套件中还包含test_sdpa_can_compile_dynamic(test_modeling_beit.py)这类对 SDPA 与torch.compile组合的回归验证。
上手资源与验证途径
- 图像分类:
BeitForImageClassification有配套的官方示例脚本 run_image_classification.py,可参照 图像分类任务指南 完成自定义数据集微调(将ViTImageProcessor/ViTForImageClassification替换为对应的 BEiT 类即可); - 语义分割:参见 语义分割任务指南;
- 模型对比:BEiT 在 ImageNet-1K 与 CIFAR-100 微调后,性能优于同为 ViT 架构的 原始 ViT 与数据高效的 DeiT;
- 测试验证:单元测试 test_modeling_beit.py 覆盖
BeitModel、四个任务模型与骨干网络的梯度、前向输出与 pipeline 集成,test_image_processing_beit.py 覆盖图像预处理与分割后处理,可通过它们校验自己的使用方式是否正确。
总而言之,BEiT 用 BERT 的掩码思想统一了视觉表示学习,而本仓库则将其落地为开箱即用的工程实现:完整的配置体系、双后端图像处理器、四类任务头加骨干网络,以及默认启用的 SDPA 加速。无论是复现论文、微调下游任务,还是将其作为检测分割框架的骨干,都可以基于上文内容直接开始。
【免费下载链接】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),仅供参考