TensorFlow Model Garden 中的 MaxViT:多轴视觉 Transformer 的架构剖析、配置详解与训练实践
2026/9/5 18:12:14 网站建设 项目流程

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)给出的执行顺序为五个子分支,每个子分支后都接一条残差连接:

  1. 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 部分;
  2. Block attention 分支block_attn_branch):先做 LayerNorm,再用window_partition把特征切成窗口,调用Attention层,最后window_stitch_back拼回原空间。源码注释明确指出这是 "local block-attention";
  3. Block FFN 分支block_ffn_branch):标准位置前馈网络(扩展率 4、GELU);
  4. Grid attention 分支grid_attn_branch):LayerNorm 后grid_partition做稀疏全局采样,再走同一个Attention层实现,最后grid_stitch_back还原;
  5. Grid FFN 分支grid_ffn_branch):第二个前馈网络。

值得注意的两个实现细节:

  • 两个注意力头共享同一套Attention实现,区别只在输入 token 的组织方式(窗口 vs 网格)。Attention层(layers.py L108-L312)基于TrailDense(einsum 实现的批量投影)构造 Q/K/V/O,默认head_size=32num_heads = hidden_size // head_size
  • 2D 相对位置偏置rel_attn_type支持2d_multi_head(默认)与2d_single_head2d_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_sizegrid_size的约束来源。

三、完整骨干MaxViT与五种规格MAXVIT_SPECS

骨干类MaxViT(maxvit.py L450-L825)由 Stem + 4 个 Stage 组成。源码 L29-L72 的MAXVIT_SPECS定义了全部五种规格(另有maxvit-tiny-for-test测试用规格):

规格survival_probstem_hsizenum_blocks (4 阶段)hidden_size (4 阶段)
maxvit-tiny0.8(64, 64)(2, 2, 5, 2)(64, 128, 256, 512)
maxvit-small0.7(64, 64)(2, 2, 5, 2)(96, 192, 384, 768)
maxvit-base0.6(64, 64)(2, 6, 14, 2)(96, 192, 384, 768)
maxvit-large0.4(128, 128)(2, 6, 14, 2)(128, 256, 512, 1024)
maxvit-xlarge0.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_sizegrid_sizescale_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_sizeNoneNone时完全采用model_name对应规格,显式设置则覆盖
head_size32每个注意力头维度,num_heads缺省时取hidden_size // head_size
rel_attn_type'2d_multi_head'可选2d_multi_head/2d_single_head/None
scale_ratioNone形如'12/7'的字符串,见下文
downsample_loc'depth_conv'MBConv 中执行下采样的位置
kernel_size3卷积核大小
se_ratio0.25SE 层瓶颈比例
data_format'channels_last'源码注明目前仅支持 channels_last
norm_type'sync_batch_norm'可选batch_norm/sync_batch_norm/layer_norm,同步 BN 适合多机训练
add_pos_encFalse是否加绝对位置编码
pool_type/pool_stride'2d:avg'/ 2下采样方式(2d:avg2d:max1d:avg1d:max)与步长
expansion_rate4MBConv 与 FFN 的扩展率
activation'gelu'激活函数
survival_prob/survival_prob_annealNone/ True随机深度存活概率;None时用规格默认值,退火使深层正则更强
representation_size/add_gap_layer_normNone/ 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-Tiny224×22483.1 (-0.5)83.631M5.6Gmaxvit_tiny_imagenet.yaml
MaxViT-Small224×22484.1 (-0.3)84.469M11.7Gmaxvit_small_imagenet.yaml
MaxViT-Base224×22484.2 (-0.7)84.9120M23.4Gmaxvit_base_imagenet.yaml
MaxViT-Large224×22484.6 (-0.6)85.2212M43.9Gmaxvit_large_imagenet.yaml
MaxViT-XLarge224×22484.8-475M97.9Gmaxvit_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-Base384×38488.37% (-0.32%)88.69%120M74.2Gfinetune_maxvitb_imagenet_i384.yaml
MaxViT-Base512×51288.63% (-0.19%)88.82%120M138.3Gfinetune_maxvitb_imagenet_i512.yaml
MaxViT-Large384×38488.86% (-0.26%)89.12%212M128.7Gfinetune_maxvitl_imagenet_i384.yaml
MaxViT-Large512×51289.02% (-0.39%)89.41%212M245.2Gfinetune_maxvitl_imagenet_i512.yaml
MaxViT-XLarge384×38489.21% (-0.15%)89.36%475M293.7Gfinetune_maxvitxl_imagenet_i384.yaml
MaxViT-XLarge512×51289.31% (-0.22%)89.53%475M535.2Gfinetune_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 预训练骨干(第一组):

模型输入尺寸窗口尺寸Epochsbox AP论文 box APmask AP配置
MaxViT-Tiny640×64020×2020049.97-42.69coco_maxvitt_i640_crcnn.yaml
MaxViT-Tiny896×89628×2820052.35 (+0.25)52.144.69-
MaxViT-Small640×64020×2020050.79-43.36-
MaxViT-Small896×89628×2820053.54 (+0.44)53.145.79coco_maxvits_i896_crcnn.yaml
MaxViT-Base640×64020×2020051.59-44.07coco_maxvitb_i640_crcnn.yaml
MaxViT-Base896×89628×2820053.47 (+0.07)53.445.96coco_maxvitb_i896_crcnn.yaml

JFT-300M 预训练骨干(第二组):

模型输入尺寸窗口尺寸Epochsbox AP论文 box APmask AP配置
MaxViT-Base896×89628×2820054.31 (+0.91)53.446.31coco_maxvitb_i896_crcnn.yaml
MaxViT-Large896×89628×2820054.69-46.59coco_maxvitl_i896_crcnn.yaml

对应的 coco_maxvitb_i896_crcnn.yaml 展示了检测场景的窗口配置:window_size: 28grid_size: 28scale_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#FLOPsglobalPR-AUC
MaxViT-Base224×224120M23.4G52.75%
MaxViT-Large224×224212M43.9G53.77%
MaxViT-XLarge224×224475M-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 中experimentmodemodel_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 提供了三个层次的验证,可用于改动源码后做回归检查:

  1. 单块前向testMaxViTBlockCreation[2, 64, 64, 3]输入构造MaxViTBlock(hidden_size=8, head_size=4, window_size=4, grid_size=4),断言输出形状[2, 64, 64, 8]且 dtype 为 float32;
  2. 整骨干前向(参数化用例):覆盖 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=16pre_logits输出形状校验;
  3. 配置化构建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_SPECSMaxViTBlockMaxViT
  • 底层算子:modeling/layers.py(AttentionFFNMBConvBlockTrailDense)、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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询