TensorFlow Model Garden 中的 MaxViT:多轴视觉 Transformer 的架构剖析、配置详解与训练实践
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
MaxViT(Multi-Axis Vision Transformer,ECCV 2022)是一族 CNN 与 ViT 混合的视觉骨干网络(backbone),当前仓库official/projects/maxvit/提供了其完整的 TensorFlow 2 实现,覆盖 ImageNet 分类预训练/微调与 COCO 检测等下游任务。本文以官方文档 MaxViT README 为骨架,结合 maxvit.py、layers.py 与 configs/backbones.py 的源码实现,讲透每个 MaxViT 块中 MBConv、block attention(窗口局部注意力)与 grid attention(膨胀全局注意力)的组成方式、关键超参(window_size/grid_size/scale_ratio)的约束与取值逻辑,以及如何复用仓库自带的实验 YAML 复现论文结果。读完本文,你将掌握从源码层面理解混合骨干网络、按任务正确配置窗口/网格参数、并运行分类与检测训练流程的完整能力。
一、MaxViT 的核心思想:混合骨干 + 线性复杂度注意力
README 对 MaxViT 的定位非常明确:它是一族hybrid (CNN + ViT) 视觉骨干模型,在参数效率(#Param)与 FLOPs 效率两个维度上,整体优于当时的 ConvNet 与 Transformer 骨干;并且能良好扩展到 ImageNet-21K 级别的大规模数据。其最关键的设计卖点是grid attention 的线性复杂度——正因为注意力复杂度对 token 数是线性的,MaxViT 才能在需要大分辨率输入的任务上(目标检测、语义分割)依然具备可扩展性。
从源码结构看,这一"线性复杂度"来自 maxvit.py 中grid_partition(源码 L310-L338)实现的稀疏/膨胀(dilated)全局注意力:
- block attention(局部):
window_partition把特征图[B, H, W, C]切分成互不重叠的窗口块,reshape 为[B·nH·nW, w, w, C],注意力只在每个w×w窗口内计算,复杂度 O(w²) 与窗口面积相关,与全图无关; - grid attention(全局):
grid_partition把特征图按 stride=grid_size的网格重排(reshape 为(-1, grid, H//grid, grid, W//grid, C)再转置),使得每个子序列内的 token 在全图上是等间隔采样的。每个子序列长度为grid_size²,注意力复杂度只与grid_size²相关,从而对 token 总数呈线性。
README 中的元架构描述与源码完全对应:"每个 MaxViT 块包含 MBConv、block attention(window-based local attention)、grid attention(dilated global attention)",整个骨干是这种块的同构堆叠(homogeneously stacked backbone)。需要说明的是,README 开头带有免责声明:该实现当时仍处于持续开发中("This implementation is still under development"),属于研究性项目代码。
二、解剖 MaxViT 块:MaxViTBlock的前向流程
核心单元是 maxvit.py 中的MaxViTBlock,其 docstring 一句话概括了组成:
"""MaxViT block = MBConv + Block-Attention + FFN + Grid-Attention + FFN."""call方法(源码 L400-L447)给出的执行顺序为五个子分支,每个子分支后都接一条残差连接:
- MBConv 分支(
mbconv_branch):Mobile Inverted Residual Bottleneck,来自 layers.py 的MBConvBlock,内部结构为 Pre-Norm → 1×1 扩展卷积(expansion_rate=4)→ 3×3 深度卷积 →SE(Squeeze-and-Excitation,se_ratio=0.25)→ 1×1 压缩卷积,是块中承担局部感知的 CNN 部分; - Block attention 分支(
block_attn_branch):先做 LayerNorm,再用window_partition把特征切成窗口,调用Attention层,最后window_stitch_back拼回原空间。源码注释明确指出这是 "local block-attention"; - Block FFN 分支(
block_ffn_branch):标准位置前馈网络(扩展率 4、GELU); - Grid attention 分支(
grid_attn_branch):LayerNorm 后grid_partition做稀疏全局采样,再走同一个Attention层实现,最后grid_stitch_back还原; - Grid FFN 分支(
grid_ffn_branch):第二个前馈网络。
值得注意的两个实现细节:
- 两个注意力头共享同一套
Attention实现,区别只在输入 token 的组织方式(窗口 vs 网格)。Attention层(layers.py L108-L312)基于TrailDense(einsum 实现的批量投影)构造 Q/K/V/O,默认head_size=32,num_heads = hidden_size // head_size; - 2D 相对位置偏置:
rel_attn_type支持2d_multi_head(默认)与2d_single_head。2d_multi_head下每个头学习一个形状为[num_heads, 2h-1, 2w-1]的可学习偏置,通过reindex_2d_einsum_lookup重索引后加到注意力 logits 上。这一机制与后文scale_ratio的微调技巧直接相关:当微调分辨率/窗口与预训练不同时,偏置词表会按scale_ratio缩小,再用tf.image.resize双线性插值回当前尺寸(layers.py L269-L283),保证位置偏置可以跨分辨率复用; - 残差连接带"存活概率":每个残差相加都经过
ops.residual_add(output, shortcut, self._survival_prob, training),即随机深度(stochastic depth)式的 DropConnect 正则化。
此外,window_partition/grid_partition都带有一个硬性校验(maxvit.py L276-L280):特征图尺寸必须能被window_size/grid_size整除,否则直接抛出ValueError。这正是配置项window_size、grid_size的约束来源。
三、完整骨干MaxViT与五种规格MAXVIT_SPECS
骨干类MaxViT(maxvit.py L450-L825)由 Stem + 4 个 Stage 组成。源码 L29-L72 的MAXVIT_SPECS定义了全部五种规格(另有maxvit-tiny-for-test测试用规格):
| 规格 | survival_prob | stem_hsize | num_blocks (4 阶段) | hidden_size (4 阶段) |
|---|---|---|---|---|
| maxvit-tiny | 0.8 | (64, 64) | (2, 2, 5, 2) | (64, 128, 256, 512) |
| maxvit-small | 0.7 | (64, 64) | (2, 2, 5, 2) | (96, 192, 384, 768) |
| maxvit-base | 0.6 | (64, 64) | (2, 6, 14, 2) | (96, 192, 384, 768) |
| maxvit-large | 0.4 | (128, 128) | (2, 6, 14, 2) | (128, 256, 512, 1024) |
| maxvit-xlarge | 0.3 | (192, 192) | (2, 6, 14, 2) | (192, 384, 768, 1536) |
可见 small 与 base 宽度相同(96/192/384/768),差异在深度(tiny/small 是 2/2/5/2,base/large/xlarge 加深为 2/6/14/2);large/xlarge 则进一步加宽。这与 README 性能表中 Tiny 31M → Small 69M → Base 120M → Large 212M → XLarge 475M 的参数增长曲线一致。
构建骨干时还有几个值得源码级关注的机制:
- Stem 只下采样 4 倍:Stem 是两个 3×3 Conv2D(首个 stride=2,第二个 stride=1),中间夹 BN 与激活(源码 L597-L619)。后续每个 Stage 的首个块以
pool_stride=2下采样(其余块 stride=1),因此总下采样率为 4/8/16/32,四个 Stage 的输出分别对应检测任务常用的 P2–P5 特征层级; - 多尺度输出端点:
call把每个 Stage 输出存入endpoints['2']…endpoints['5'],并通过output_specs属性暴露形状(源码 L799-L825)——这是它能直接对接 RetinaNet / Cascade RCNN / 分割 head 的关键接口。maxvit_test.py 中testBuildMaxViTWithConfig正是断言output_specs的键集合为{'2','3','4','5'}; - 分类头是可选的:仅当
representation_size > 0时,骨干末尾才接 GlobalAveragePooling2D →(可选 LayerNorm,add_gap_layer_norm)→ Dense →tanh,输出pre_logits(源码 L812-L819); - 随机深度退火:当
survival_prob_anneal=True(默认),每个块的存活概率从 1.0 按块序线性退火到规格给定的survival_prob(源码 L638-L649),即浅层几乎不丢弃、深层正则更强; - 绝对位置编码默认关闭(
add_pos_enc=False)。MaxViT 依靠 CNN 的局部归纳偏置 + 2D 相对位置偏置定位,可选的绝对正弦位置编码只加在第三个 Stage 的首个块输入上(源码 L740-L764)。
骨干的注册与构建入口
build_maxvit通过装饰器@factory.register_backbone_builder('maxvit')注册进 Model Garden 的骨干工厂(maxvit.py L915-L933),并会用假输入前向一次以获得正确的output_specs。构建逻辑override_predefined_spec_and_build_maxvit的规则是:先用MAXVIT_SPECS[model_name]取默认规格,再被 config 中显式设置的stem_hsize/block_type/num_blocks/hidden_size逐项覆盖。
四、关键配置项详解:window_size、grid_size与scale_ratio
MaxViT 的配置定义在 configs/backbones.py 中,MaxViT是一个hyperparams.Configdataclass。其中最核心、也最容易被配错的参数是窗口/网格尺寸,源码注释给出了非常实用的经验法则(backbones.py L36-L44):
# Note that the window_size and grid_size should be divisible by all the # feature map sizes along the entire network. Say, if you train on ImageNet # classification at 224x224, set both to 7 is almost the only choice. # If you train on COCO object detection at 896x896, set it to 28 is suggested, # as following Swin Transformer, window size should scales with feature size. # You may as well set it as 14 or 7. window_size: int = 7 # window size for conducting block attention module. grid_size: int = 7 # grid size for conducting sparse global grid attention.结合骨干结构可以推导这条规则:224×224 输入下四个 Stage 的特征图是 56/28/14/7,7 是其中唯一的公共约数,所以 224 训练时window_size=grid_size=7几乎是唯一选择;而 COCO 896×896 的特征图是 224/112/56/28,按 Swin Transformer 的思路"窗口尺寸随特征尺寸等比放大",建议取 28(也可以取 14 或 7)。
其他关键配置项及其默认值(backbones.py L26-L87):
| 配置项 | 默认值 | 说明 |
|---|---|---|
model_name | 'maxvit-tiny' | 选择MAXVIT_SPECS中的预定义规格 |
stem_hsize/block_type/num_blocks/hidden_size | None | 为None时完全采用model_name对应规格,显式设置则覆盖 |
head_size | 32 | 每个注意力头维度,num_heads缺省时取hidden_size // head_size |
rel_attn_type | '2d_multi_head' | 可选2d_multi_head/2d_single_head/None |
scale_ratio | None | 形如'12/7'的字符串,见下文 |
downsample_loc | 'depth_conv' | MBConv 中执行下采样的位置 |
kernel_size | 3 | 卷积核大小 |
se_ratio | 0.25 | SE 层瓶颈比例 |
data_format | 'channels_last' | 源码注明目前仅支持 channels_last |
norm_type | 'sync_batch_norm' | 可选batch_norm/sync_batch_norm/layer_norm,同步 BN 适合多机训练 |
add_pos_enc | False | 是否加绝对位置编码 |
pool_type/pool_stride | '2d:avg'/ 2 | 下采样方式(2d:avg、2d:max、1d:avg、1d:max)与步长 |
expansion_rate | 4 | MBConv 与 FFN 的扩展率 |
activation | 'gelu' | 激活函数 |
survival_prob/survival_prob_anneal | None/ True | 随机深度存活概率;None时用规格默认值,退火使深层正则更强 |
representation_size/add_gap_layer_norm | None/ True | 分类头宽度(须与最后一个 Stage 的hidden_size一致)与 GAP 后 LayerNorm |
scale_ratio是跨分辨率/跨窗口微调的开关:它记录"当前窗口尺寸 / checkpoint 窗口尺寸",用于把预训练学到的 2D 相对位置偏置按词表缩小后双线性插值回当前尺寸(实现见上文 layers.py L269-L283)。仓库里的实验 YAML 正是按这一机制成套设置的:ImageNet 224 预训练用 7,384 微调用 12('12/7'),COCO 896 检测用 28('28/7')。
五、实验配置与性能结果
README 给出的结果分四组,此处完整继承并配齐对应配置路径(路径均相对仓库根目录)。
注意(README 原文说明):DeiT ImageNet 预训练的实验设置与论文不同——这里遵循论文的预训练超参、仅跑相近的训练步数,而论文建议以不同超参 + EMA 做短程微调,因此表中数字会比论文值略低(表中括号内为与论文值的差)。
5.1 DeiT 风格 ImageNet-1k 预训练
| 模型 | 评测尺寸 | Top-1 Acc | 论文 Acc | #Param | #FLOPs | 配置 |
|---|---|---|---|---|---|---|
| MaxViT-Tiny | 224×224 | 83.1 (-0.5) | 83.6 | 31M | 5.6G | maxvit_tiny_imagenet.yaml |
| MaxViT-Small | 224×224 | 84.1 (-0.3) | 84.4 | 69M | 11.7G | maxvit_small_imagenet.yaml |
| MaxViT-Base | 224×224 | 84.2 (-0.7) | 84.9 | 120M | 23.4G | maxvit_base_imagenet.yaml |
| MaxViT-Large | 224×224 | 84.6 (-0.6) | 85.2 | 212M | 43.9G | maxvit_large_imagenet.yaml |
| MaxViT-XLarge | 224×224 | 84.8 | - | 475M | 97.9G | maxvit_xlarge_imagenet.yaml |
以 maxvit_base_imagenet.yaml 为例,预训练的核心训练超参为:AdamW(weight_decay_rate: 0.05)、EMA(average_decay: 0.9999)、cosine 学习率(初始 0.003、alpha: 0.01)、线性 warmup 10000 步(从 0 起步)。该目录下还有各规格的_gpu.yaml变体,便于非 TPU 环境使用。
5.2 ImageNet 预训练权重的微调(大分辨率)
| 模型 | 输入尺寸 | Top-1 Acc | 论文 Acc | #Param | #FLOPs | 配置 |
|---|---|---|---|---|---|---|
| MaxViT-Base | 384×384 | 88.37% (-0.32%) | 88.69% | 120M | 74.2G | finetune_maxvitb_imagenet_i384.yaml |
| MaxViT-Base | 512×512 | 88.63% (-0.19%) | 88.82% | 120M | 138.3G | finetune_maxvitb_imagenet_i512.yaml |
| MaxViT-Large | 384×384 | 88.86% (-0.26%) | 89.12% | 212M | 128.7G | finetune_maxvitl_imagenet_i384.yaml |
| MaxViT-Large | 512×512 | 89.02% (-0.39%) | 89.41% | 212M | 245.2G | finetune_maxvitl_imagenet_i512.yaml |
| MaxViT-XLarge | 384×384 | 89.21% (-0.15%) | 89.36% | 475M | 293.7G | finetune_maxvitxl_imagenet_i384.yaml |
| MaxViT-XLarge | 512×512 | 89.31% (-0.22%) | 89.53% | 475M | 535.2G | finetune_maxvitxl_imagenet_i512.yaml |
以 finetune_maxvitb_imagenet_i384.yaml 为例,微调配方体现了与预训练完全不同的超参哲学,并演示了前文讲的窗口缩放机制:
runtime: mixed_precision_dtype: 'bfloat16' # bfloat16 混合精度 task: init_checkpoint: 'Please provide' # 224 预训练 checkpoint 路径(需自行提供) init_checkpoint_modules: 'backbone' # 仅恢复骨干权重 model: backbone: maxvit: model_name: 'maxvit-base' representation_size: 768 survival_prob: 0.8 # 微调时提高存活概率,减弱正则 window_size: 12 grid_size: 12 scale_ratio: '12/7' # 384 分辨率下窗口 12 = 7 × (384/224) input_size: [384, 384, 3] train_data: global_batch_size: 512 aug_type: { type: 'randaug', randaug: { magnitude: 15 } } losses: label_smoothing: 0.1 trainer: train_steps: 100080 optimizer_config: optimizer: type: 'adamw' adamw: { weight_decay_rate: 1.0e-4, gradient_clip_norm: 1.0 } ema: { average_decay: 0.9999, trainable_weights_only: false } learning_rate: { type: constant, constant: { learning_rate: 5.0e-5 } } warmup: { type: null }对比 224 预训练配置可见:微调把学习率从 cosine 0.003 降到常数 5e-5(无 warmup)、权重衰减从 0.05 降到 1e-4 并加了梯度裁剪、survival_prob提到 0.8、窗口/网格从 7 升到 12 并用scale_ratio: '12/7'复用 224 checkpoint 的位置偏置——这套参数是复现表中高分的关键。
5.3 COCO Cascade RCNN(检测/分割)
DeiT 预训练骨干(第一组):
| 模型 | 输入尺寸 | 窗口尺寸 | Epochs | box AP | 论文 box AP | mask AP | 配置 |
|---|---|---|---|---|---|---|---|
| MaxViT-Tiny | 640×640 | 20×20 | 200 | 49.97 | - | 42.69 | coco_maxvitt_i640_crcnn.yaml |
| MaxViT-Tiny | 896×896 | 28×28 | 200 | 52.35 (+0.25) | 52.1 | 44.69 | - |
| MaxViT-Small | 640×640 | 20×20 | 200 | 50.79 | - | 43.36 | - |
| MaxViT-Small | 896×896 | 28×28 | 200 | 53.54 (+0.44) | 53.1 | 45.79 | coco_maxvits_i896_crcnn.yaml |
| MaxViT-Base | 640×640 | 20×20 | 200 | 51.59 | - | 44.07 | coco_maxvitb_i640_crcnn.yaml |
| MaxViT-Base | 896×896 | 28×28 | 200 | 53.47 (+0.07) | 53.4 | 45.96 | coco_maxvitb_i896_crcnn.yaml |
JFT-300M 预训练骨干(第二组):
| 模型 | 输入尺寸 | 窗口尺寸 | Epochs | box AP | 论文 box AP | mask AP | 配置 |
|---|---|---|---|---|---|---|---|
| MaxViT-Base | 896×896 | 28×28 | 200 | 54.31 (+0.91) | 53.4 | 46.31 | coco_maxvitb_i896_crcnn.yaml |
| MaxViT-Large | 896×896 | 28×28 | 200 | 54.69 | - | 46.59 | coco_maxvitl_i896_crcnn.yaml |
对应的 coco_maxvitb_i896_crcnn.yaml 展示了检测场景的窗口配置:window_size: 28、grid_size: 28、scale_ratio: '28/7'(与 4.4 节的 896 规则一致),survival_prob: 0.2(检测任务正则更强),init_checkpoint_modules: ['backbone']只恢复骨干,其余训练超参为 AdamW(wd 0.05)+ EMA 0.9998 + cosine 0.003 学习率、6000 步 warmup、90000 步(对应 200 epoch,global_batch_size: 256)、l2_weight_decay: 2.0e-07。此外 configs/experiments 下还提供 RetinaNet(retinanet_maxvit_base_coco_i640_tpu.yaml)与语义分割(seg_coco_maxvits_i640.yaml、seg_pascal_maxvits_i512.yaml)配置,印证了 README "在检测与分割任务上良好扩展" 的说法。
5.4 JFT-300M 监督式预训练(下游 globalPR-AUC)
| 模型 | 预训练尺寸 | #Param | #FLOPs | globalPR-AUC |
|---|---|---|---|---|
| MaxViT-Base | 224×224 | 120M | 23.4G | 52.75% |
| MaxViT-Large | 224×224 | 212M | 43.9G | 53.77% |
| MaxViT-XLarge | 224×224 | 475M | - | 54.71% |
六、运行训练:入口、参数与配置装配
MaxViT 项目的训练入口是 train.py,内容非常精简:
"""TensorFlow Model Garden Vision training driver, including MaxViT configs..""" from absl import app from official.common import flags as tfm_flags from official.projects.maxvit import registry_imports # pylint: disable=unused-import from official.vision import train if __name__ == '__main__': tfm_flags.define_flags() app.run(train.main)从源码结构看,装配流程是:registry_imports.py导入official.vision.registry_imports与 项目 configs 包、maxvit 模块,从而把'maxvit'骨干及分类/检测/分割任务配置注册进全局工厂;随后official.vision.train.main按统一的 Vision 训练驱动器执行。official/common/flags.py 中experiment、mode、model_dir三个 flag 被标记为必填,--experiment的值即对应configs/experiments/下某个 YAML 的文件名(不含扩展名)。据此,一次典型的训练命令形如:
python official/projects/maxvit/train.py \ --experiment maxvit_base_imagenet \ --mode train \ --model_dir /path/to/checkpoints \ --dataset_dir /path/to/imagenet检测任务则把--experiment换成coco_maxvitb_i896_crcnn等即可。对于微调与检测类配置,YAML 里的init_checkpoint: 'Please provide'是占位符,需要自行提供 ImageNet-1k 或 JFT 预训练 checkpoint 的路径。运行环境依赖见仓库根目录的 requirements.txt。
七、正确性验证:maxvit_test.py的测试断言
modeling/maxvit_test.py 提供了三个层次的验证,可用于改动源码后做回归检查:
- 单块前向:
testMaxViTBlockCreation用[2, 64, 64, 3]输入构造MaxViTBlock(hidden_size=8, head_size=4, window_size=4, grid_size=4),断言输出形状[2, 64, 64, 8]且 dtype 为 float32; - 整骨干前向(参数化用例):覆盖 3 阶段/4 阶段规格、Tiny 规格(
stem_hsize=[64,64]、num_blocks=[2,3,5,2]、hidden_size=[96,192,384,768],期望最深特征为[2, 2, 2, 768],与 Stem 4 倍 + 每 Stage 2 倍下采样的推导一致),以及带representation_size=16的pre_logits输出形状校验; - 配置化构建:
testBuildMaxViTWithConfig验证经由backbones.Backbone(type='maxvit')+build_maxvit的注册路径可用,并断言output_specs键为{'2','3','4','5'}。
八、引用信息
README 建议引用原文时使用的 BibTeX(作者为 Tu, Zhengzhong; Talebi, Hossein; Zhang, Han; Yang, Feng; Milanfar, Peyman; Bovik, Alan; Li, Yinxiao,发表于 ECCV 2022):
@article{tu2022maxvit, title={MaxViT: Multi-Axis Vision Transformer}, author={Tu, Zhengzhong and Talebi, Hossein and Zhang, Han and Yang, Feng and Milanfar, Peyman and Bovik, Alan and Li, Yinxiao}, journal={ECCV}, year={2022}, }延伸阅读路径
- 架构与结果总览:official/projects/maxvit/README.md
- 骨干实现:modeling/maxvit.py(
MAXVIT_SPECS、MaxViTBlock、MaxViT) - 底层算子:modeling/layers.py(
Attention、FFN、MBConvBlock、TrailDense)、modeling/common_ops.py - 配置定义:configs/backbones.py,各任务配置位于 configs/experiments/
- 训练入口与注册:train.py、registry_imports.py
- 单元测试:modeling/maxvit_test.py
【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考