FocalNet 模型深度解析:Transformers 中的 Focal Modulation Networks 视觉骨干网络
【免费下载链接】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
FocalNet(Focal Modulation Networks)是微软研究院提出的纯卷积式视觉骨干网络,用"焦点调制"(focal modulation)机制彻底替代了 ViT 与 Swin 等模型依赖的自注意力(self-attention),在图像分类、目标检测与语义分割上取得了与自注意力模型相当甚至更优的结果。本文基于本仓库中 FocalNet 官方文档 与其完整 PyTorch 实现,为你系统讲解 FocalNet 的架构原理、FocalNetConfig全部配置参数、四大模型类(含掩码图像建模与分类头)的正确用法,以及如何借助Auto*API 直接加载预训练权重并接入下游任务。
一、FocalNet 是什么:用焦点调制取代自注意力
FocalNet 由 Jianwei Yang、Chunyuan Li、Xiyang Dai、Lu Yuan、Jianfeng Gao 在论文Focal Modulation Networks中提出(论文于 2022-03-22 收录于 HF Papers,模型实现于 2023-04-23 贡献给 Transformers)。该模型的核心理念是:视觉任务中的 token 交互不一定需要自注意力,可以用一种名为 focal modulation 的机制来完成。论文摘要指出,作者用一摞深度可分离卷积层堆叠做层次化上下文编码,并辅以门控聚合与逐元素调制,从而以相近计算成本在图像分类、检测、分割上超越 Swin、Focal Transformer 等 SOTA 自注意力模型。
需要说明:论文中提到的诸如 "tiny/base 尺寸在 ImageNet-1K 上取得 82.3%/83.9% top-1 准确率"、"COCO 目标检测相较 Swin 提升 2.1 点" 等数字均出自论文摘要(收录于 focalnet.md),可作为背景参考。
原文档将其与 ViT、Swin 对照——后两者都把自注意力作为建模 token 交互的核心算子,而 FocalNet 中该角色完全由 focal modulation 承担。模块由 nielsr 贡献,权重转换基于微软官方代码。
焦点调制的三个组成部分
Focal modulation 由三部分构成,在源码 modeling_focalnet.py 的FocalNetModulation(第 245-313 行)中一一对应:
- 层次化上下文编码(hierarchical contextualization):通过一组逐层卷积核放大的 depth-wise 卷积把视觉上下文从"短程"聚合到"长程"。源码中对应
self.focal_layers这个ModuleList,第k层的卷积核大小为focal_factor * k + focal_window(focal_factor=2),每个卷积后都接 GELU 激活; - 门控聚合(gated aggregation):为每个查询 token 依据其内容选择性聚合不同范围的上下文。源码中线性投影
projection_in把输入拆成q、ctx、gates三份,其中gates有focal_level + 1个通道,前focal_level个通道分别加权不同 focal 层的局部上下文,最后一个通道加权全局平均池化后的上下文(ctx_global); - 逐元素调制 / 仿射变换(element-wise modulation):把聚合后的上下文调制到查询上。源码中
ctx_all经 1x1 卷积projection_context得到modulator,与查询q逐元素相乘(x_out = q * modulator),最后再经projection_out输出。
因此,一个FocalNetLayer的残差块结构为:LayerNorm → FocalNetModulation → DropPath → 残差相加 → MLP → 残差相加。这与 Swin 的 block 结构非常相似(源码中FocalNetDropPath、FocalNetForImageClassification等甚至直接标注 "Copied from transformers.models.swin..."),区别仅在于把窗口注意力换成了 focal modulation——这正是它"零注意力"纯卷积设计的核心。
二、网络整体结构与数据流
FocalNet 主体是一个 4 阶段(4-stage)分层金字塔编码器,实现于 modeling_focalnet.py:
- Patch 嵌入(stem):
FocalNetPatchEmbeddings(第 174-242 行)用patch_size=4、stride=4 的卷积把图像切成 patch 序列(token 数num_patches = (H/patch_h) * (W/patch_w)),输入尺寸不能整除 patch 时会被maybe_pad自动补齐;token 维度即embed_dim。 - 4 个 stage:
FocalNetStage(第 429-492 行)由depths[i]个FocalNetLayer组成。相邻 stage 之间用一个downsample(本质上仍是FocalNetPatchEmbeddings,patch_size=2、stride=2)把空间分辨率减半、通道数翻倍。token 数从224/4=56依次降到28 → 14 → 7,对应总下采样率 32。 - Stochastic depth:各 layer 的 drop path 率按
torch.linspace(0, config.drop_path_rate, sum(depths))规则从 0 线性递增到drop_path_rate,越深的层丢弃概率越高。 - LayerScale 可选:
use_layerscale=True时每个 block 的gamma_1/gamma_2初始化为layerscale_value(默认 1e-4)并可学习,见FocalNetPreTrainedModel._init_weights。
FocalNetModel.forward接收pixel_values(形状(batch, channels, H, W)),输出FocalNetModelOutput,其中:
last_hidden_state:(batch, seq_len, hidden_size)的序列输出;pooler_output:可选的AdaptiveAvgPool1d平均池化结果;hidden_states与reshaped_hidden_states:当output_hidden_states=True时返回各 stage 的输出,后者被 reshape 回带空间维度的(batch, hidden_size, height, width)形式,方便下游密集预测任务直接消费。
三、FocalNetConfig 配置参数详解
FocalNetConfig继承自BackboneConfigMixin与PreTrainedConfig,源码见 configuration_focalnet.py。除标准 ViT 类模型的hidden_act、mlp_ratio、drop_path_rate、initializer_range、layer_norm_eps、num_labels等常规参数外,其 FocalNet 特有的核心参数如下表(含默认值与含义):
| 参数 | 默认值 | 说明 |
|---|---|---|
image_size | 224 | 输入图像分辨率(可为 int 或 tuple,见FocalNetPatchEmbeddings) |
patch_size | 4 | patch 边长,stem 卷积的核/步长 |
num_channels | 3 | 输入图像通道数 |
embed_dim | 96 | stem 输出通道数(第一维特征维) |
use_conv_embed | False | 是否使用卷积嵌入。作者指出使用卷积嵌入通常能提升性能,但默认不开(stem 用 7x7 卷积、内部 stage 用 3x3,见源码第 196-210 行);开启后实测模型多为带lrf(large receptive field)后缀的权重 |
hidden_sizes | (192, 384, 768, 768) | 各 stage 的输出维度(主要用于 backbone 的num_features) |
depths | (2, 2, 6, 2) | 各 stage 的 block 数量 |
focal_levels | (2, 2, 2, 2) | 各 stage 中 focal modulation 的焦点层级数(focal layer 卷积个数) |
focal_windows | (3, 3, 3, 3) | 各 stage 的 focal window 大小(最小感受野基线) |
mlp_ratio | 4.0 | MLP 隐藏层维数 =dim * mlp_ratio |
hidden_dropout_prob | 0.0 | embedding 与全连接层的 dropout 概率 |
drop_path_rate | 0.1 | stochastic depth 最大概率 |
use_layerscale | False | 是否在残差块中使用 LayerScale |
layerscale_value | 0.0001 | LayerScale 的初始值 |
use_post_layernorm | False | 是否使用 post-LayerNorm(否则为 pre-LN,见FocalNetLayer第 414-417 行) |
use_post_layernorm_in_modulation | False | focal modulation 内部是否使用 post-LayerNorm |
normalize_modulator | False | 是否归一化 modulator(开启后上下文除以focal_level + 1) |
encoder_stride | 32 | MIM decoder head 需要的上采样倍数 |
out_features/out_indices | None | Backbone 要输出的 stage 集合/索引,见FocalNetBackbone |
__post_init__会自动生成stage_names = ["stem", "stage1", "stage2", "stage3", "stage4"],供 backbone 按名取特征。
>>> from transformers import FocalNetConfig, FocalNetModel >>> # 初始化一个 microsoft/focalnet-tiny 风格配置 >>> configuration = FocalNetConfig() >>> # 由配置初始化随机权重模型 >>> model = FocalNetModel(configuration) >>> # 访问模型配置 >>> configuration = model.config四、模型类族与 forward 接口
原文档中列出的类定义可在init.py 的导入结构中确认,全部类名见源码文件末尾__all__:
1.FocalNetModel
最底层的编码器模型。参数:pixel_values(必填)、bool_masked_pos(MIM 掩码)、output_hidden_states、return_dict。返回FocalNetModelOutput。FocalNetEmbeddings内含一个可选的mask_token(仅当use_mask_token=True时创建),掩码替换逻辑见第 164-168 行。
2.FocalNetForMaskedImageModeling
在编码器之上加了 MIM decoder,实现与 SimMIM 一致(文档注明参见 SimMIM 论文)。decoder 由 1x1 卷积 +PixelShuffle(encoder_stride)组成(第 690-695 行),把编码器最后一层特征直接上采样回原始分辨率并重建像素。配合bool_masked_pos计算掩码区域的 L1 重建损失(第 768-769 行)。损失只统计被 mask 的 patch 位置,并在每个 patch 内部做repeat_interleave展开回像素粒度。
完整用法示例(取自该类的 docstring):
>>> from transformers import AutoImageProcessor, FocalNetConfig, FocalNetForMaskedImageModeling >>> 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("microsoft/focalnet-base-simmim-window6-192") >>> config = FocalNetConfig() >>> model = FocalNetForMaskedImageModeling(config) >>> num_patches = (model.config.image_size // model.config.patch_size) ** 2 >>> pixel_values = image_processor(images=image, return_tensors="pt").pixel_values >>> # 生成形状为 (batch_size, 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, reconstructed_pixel_values = outputs.loss, outputs.logits >>> list(reconstructed_pixel_values.shape) [1, 3, 192, 192]运行前请安装本仓库 examples/pytorch/image-pretraining 目录所需的依赖;文档同时提示该目录提供了用于在自定义数据上预训练 MIM 模型的脚本。注意:MIM 训练通常需把encoder_stride(默认 32)与image_size/patch_size匹配,不同 SimMIM 权重可能使用不同的分辨率与 window(如上述 192 分辨率权重)。
3.FocalNetForImageClassification
在FocalNetModel(默认开启池化层)之上接一个线性分类头(num_labels > 0时为nn.Linear(num_features, num_labels),否则退化为Identity),对池化输出分类。num_labels=1时计算 MSE 回归损失,num_labels>1时计算交叉熵。返回FocalNetImageClassifierOutput(含loss、logits)。
>>> from transformers import AutoImageProcessor, AutoModelForImageClassification >>> 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())) >>> processor = AutoImageProcessor.from_pretrained("microsoft/focalnet-tiny") >>> model = AutoModelForImageClassification.from_pretrained("microsoft/focalnet-tiny") >>> inputs = processor(images=image, return_tensors="pt") >>> outputs = model(**inputs) >>> logits = outputs.logits >>> predicted_label = logits.argmax(-1).item() # 结合 processor.id2label 得到类别名4.FocalNetBackbone(源码补充)
原文档的 autodoc 列表虽未单列,但 modeling_focalnet.py 中实现了FocalNetBackbone,文档注释表明其面向 X-Decoder 等多尺度特征消费框架。它通过配置中的out_features/out_indices(如"stage1"..."stage4"或索引)从reshaped_hidden_states中筛选输出特征金字塔feature_maps。官方权重大多为带L 形大感受野(L shape / LRF)的下采样卷积设计,加载示例(取自类 docstring):
>>> from transformers import AutoImageProcessor, AutoBackbone >>> 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())) >>> processor = AutoImageProcessor.from_pretrained("microsoft/focalnet-tiny-lrf") >>> model = AutoBackbone.from_pretrained("microsoft/focalnet-tiny-lrf") >>> inputs = processor(image, return_tensors="pt") >>> outputs = model(**inputs)五、Auto API 注册与图像处理器
FocalNet 已完整接入 Auto 体系,相关注册可在以下文件确认:
- modeling_auto.py:
AutoModel→FocalNetModel; - modeling_auto.py:
AutoModelForMaskedImageModeling→FocalNetForMaskedImageModeling; - modeling_auto.py:
AutoModelForImageClassification→FocalNetForImageClassification; - modeling_auto.py:
AutoBackbone→FocalNetBackbone; - auto_mappings.py 与 image_processing_auto.py:FocalNet 的图像处理器被映射为
BitImageProcessor(沿用 BiT 的归一化与 resize 策略),因此AutoImageProcessor.from_pretrained("microsoft/focalnet-*")即可得到正确的预处理配置。
实践中最省事的三步走就是:AutoImageProcessor预处理图像 →AutoModelForImageClassification/AutoModelForMaskedImageModeling/AutoBackbone加载权重 → 前向得到logits/reconstruction/feature_maps。如果你有自定义 checkpoint 需要转换,可以参考仓库提供的 convert_focalnet_to_hf_format.py。
六、验证与使用注意点
- 测试证据:模型行为由 tests/models/focalnet/test_modeling_focalnet.py 覆盖。
FocalNetModelTester采用image_size=32, patch_size=2, depths=[1,2,1]等小型配置做前向/梯度检查;FocalNetModelIntegrationTest则使用真实权重microsoft/focalnet-tiny(AutoImageProcessor+FocalNetForImageClassification)做端到端集成验证,并混入BackboneTesterMixin校验 backbone 输出。 - 前向输入:模型
main_input_name = "pixel_values",即 pipeline/Trainer 默认喂入名为pixel_values的张量;pixel_values缺失会直接抛错(源码第 635-636 行)。 - 显存优化:
FocalNetPreTrainedModel声明supports_gradient_checkpointing = True,配合_no_split_modules = ["FocalNetStage"],可安全启用梯度检查点以降低大输入显存占用。 - hidden states 形状:由于分层下采样,各 stage 的 hidden state 空间尺寸不同,若要拼接多尺度特征,优先使用已经 reshape 好的
reshaped_hidden_states(FocalNetEncoder内完成b(hw)c → bchw重排)。 - 预训练范围:仓库本身不包含权重文件,以上所有 checkpoint 名(
microsoft/focalnet-tiny、microsoft/focalnet-base-simmim-window6-192、microsoft/focalnet-tiny-lrf)均出自源码 docstring,实际加载需联网访问 Hugging Face Hub 对应仓库;本地只读使用时请勿尝试向仓库写入内容。
总而言之,FocalNet 为 Transformer 系视觉模型提供了一个"无需注意力"的高性价比替代方案,在 Transformers 库中其实现完整、API 清晰:理解 focal modulation 的"层级卷积 + 门控聚合 + 元素调制"三段式设计,再掌握FocalNetConfig的参数与四大类的前向约定,你就能在分类、掩码图像建模(SimMIM 式自监督预训练)与多尺度密集预测任务中自由驾驭这一骨干网络。
【免费下载链接】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),仅供参考