TensorFlow Model Garden 中的 MaxViT:多轴视觉 Transformer 的架构剖析、配置详解与训练实践

发布时间:2026/9/5 18:12:15
TensorFlow Model Garden 中的 MaxViT:多轴视觉 Transformer 的架构剖析、配置详解与训练实践 TensorFlow Model Garden 中的 MaxViT多轴视觉 Transformer 的架构剖析、配置详解与训练实践【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/modelsMaxViTMulti-Axis Vision TransformerECCV 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把特征图按 stridegrid_size的网格重排reshape 为(-1, grid, H//grid, grid, W//grid, C)再转置使得每个子序列内的 token 在全图上是等间隔采样的。每个子序列长度为grid_size²注意力复杂度只与grid_size²相关从而对 token 总数呈线性。README 中的元架构描述与源码完全对应每个 MaxViT 块包含 MBConv、block attentionwindow-based local attention、grid attentiondilated 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_branchMobile Inverted Residual Bottleneck来自 layers.py 的MBConvBlock内部结构为 Pre-Norm → 1×1 扩展卷积expansion_rate4→ 3×3 深度卷积 →SESqueeze-and-Excitationse_ratio0.25→ 1×1 压缩卷积是块中承担局部感知的 CNN 部分Block attention 分支block_attn_branch先做 LayerNorm再用window_partition把特征切成窗口调用Attention层最后window_stitch_back拼回原空间。源码注释明确指出这是 local block-attentionBlock FFN 分支block_ffn_branch标准位置前馈网络扩展率 4、GELUGrid attention 分支grid_attn_branchLayerNorm 后grid_partition做稀疏全局采样再走同一个Attention层实现最后grid_stitch_back还原Grid FFN 分支grid_ffn_branch第二个前馈网络。值得注意的两个实现细节两个注意力头共享同一套Attention实现区别只在输入 token 的组织方式窗口 vs 网格。Attention层layers.py L108-L312基于TrailDenseeinsum 实现的批量投影构造 Q/K/V/O默认head_size32num_heads hidden_size // head_size2D 相对位置偏置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骨干类MaxViTmaxvit.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/2base/large/xlarge 加深为 2/6/14/2large/xlarge 则进一步加宽。这与 README 性能表中 Tiny 31M → Small 69M → Base 120M → Large 212M → XLarge 475M 的参数增长曲线一致。构建骨干时还有几个值得源码级关注的机制Stem 只下采样 4 倍Stem 是两个 3×3 Conv2D首个 stride2第二个 stride1中间夹 BN 与激活源码 L597-L619。后续每个 Stage 的首个块以pool_stride2下采样其余块 stride1因此总下采样率为 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 →可选 LayerNormadd_gap_layer_norm→ Dense →tanh输出pre_logits源码 L812-L819随机深度退火当survival_prob_annealTrue默认每个块的存活概率从 1.0 按块序线性退火到规格给定的survival_prob源码 L638-L649即浅层几乎不丢弃、深层正则更强绝对位置编码默认关闭add_pos_encFalse。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_ratioMaxViT 的配置定义在 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/77 是其中唯一的公共约数所以 224 训练时window_sizegrid_size7几乎是唯一选择而 COCO 896×896 的特征图是 224/112/56/28按 Swin Transformer 的思路窗口尺寸随特征尺寸等比放大建议取 28也可以取 14 或 7。其他关键配置项及其默认值backbones.py L26-L87配置项默认值说明model_namemaxvit-tiny选择MAXVIT_SPECS中的预定义规格stem_hsize/block_type/num_blocks/hidden_sizeNone为None时完全采用model_name对应规格显式设置则覆盖head_size32每个注意力头维度num_heads缺省时取hidden_size // head_sizerel_attn_type2d_multi_head可选2d_multi_head/2d_single_head/Nonescale_ratioNone形如12/7的字符串见下文downsample_locdepth_convMBConv 中执行下采样的位置kernel_size3卷积核大小se_ratio0.25SE 层瓶颈比例data_formatchannels_last源码注明目前仅支持 channels_lastnorm_typesync_batch_norm可选batch_norm/sync_batch_norm/layer_norm同步 BN 适合多机训练add_pos_encFalse是否加绝对位置编码pool_type/pool_stride2d:avg/ 2下采样方式2d:avg、2d:max、1d:avg、1d:max与步长expansion_rate4MBConv 与 FFN 的扩展率activationgelu激活函数survival_prob/survival_prob_annealNone/ True随机深度存活概率None时用规格默认值退火使深层正则更强representation_size/add_gap_layer_normNone/ True分类头宽度须与最后一个 Stage 的hidden_size一致与 GAP 后 LayerNormscale_ratio是跨分辨率/跨窗口微调的开关它记录当前窗口尺寸 / checkpoint 窗口尺寸用于把预训练学到的 2D 相对位置偏置按词表缩小后双线性插值回当前尺寸实现见上文 layers.py L269-L283。仓库里的实验 YAML 正是按这一机制成套设置的ImageNet 224 预训练用 7384 微调用 1212/7COCO 896 检测用 2828/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.yamlMaxViT-Small224×22484.1 (-0.3)84.469M11.7Gmaxvit_small_imagenet.yamlMaxViT-Base224×22484.2 (-0.7)84.9120M23.4Gmaxvit_base_imagenet.yamlMaxViT-Large224×22484.6 (-0.6)85.2212M43.9Gmaxvit_large_imagenet.yamlMaxViT-XLarge224×22484.8-475M97.9Gmaxvit_xlarge_imagenet.yaml以 maxvit_base_imagenet.yaml 为例预训练的核心训练超参为AdamWweight_decay_rate: 0.05、EMAaverage_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.yamlMaxViT-Base512×51288.63% (-0.19%)88.82%120M138.3Gfinetune_maxvitb_imagenet_i512.yamlMaxViT-Large384×38488.86% (-0.26%)89.12%212M128.7Gfinetune_maxvitl_imagenet_i384.yamlMaxViT-Large512×51289.02% (-0.39%)89.41%212M245.2Gfinetune_maxvitl_imagenet_i512.yamlMaxViT-XLarge384×38489.21% (-0.15%)89.36%475M293.7Gfinetune_maxvitxl_imagenet_i384.yamlMaxViT-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.yamlMaxViT-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.yamlMaxViT-Base640×64020×2020051.59-44.07coco_maxvitb_i640_crcnn.yamlMaxViT-Base896×89628×2820053.47 (0.07)53.445.96coco_maxvitb_i896_crcnn.yamlJFT-300M 预训练骨干第二组模型输入尺寸窗口尺寸Epochsbox AP论文 box APmask AP配置MaxViT-Base896×89628×2820054.31 (0.91)53.446.31coco_maxvitb_i896_crcnn.yamlMaxViT-Large896×89628×2820054.69-46.59coco_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]只恢复骨干其余训练超参为 AdamWwd 0.05 EMA 0.9998 cosine 0.003 学习率、6000 步 warmup、90000 步对应 200 epochglobal_batch_size: 256、l2_weight_decay: 2.0e-07。此外 configs/experiments 下还提供 RetinaNetretinanet_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-AUCMaxViT-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: disableunused-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_size8, head_size4, window_size4, grid_size4)断言输出形状[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_size16的pre_logits输出形状校验配置化构建testBuildMaxViTWithConfig验证经由backbones.Backbone(typemaxvit)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 2022article{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.pyMAXVIT_SPECS、MaxViTBlock、MaxViT底层算子modeling/layers.pyAttention、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),仅供参考