如何在ImageNet上从零训练CSWin Transformer:main.py分布式训练、EMA与混合精度全流程解读

发布时间:2026/8/22 14:11:47
如何在ImageNet上从零训练CSWin Transformer:main.py分布式训练、EMA与混合精度全流程解读 如何在ImageNet上从零训练CSWin Transformermain.py分布式训练、EMA与混合精度全流程解读【免费下载链接】CSWin-TransformerCSWin Transformer: A General Vision Transformer Backbone with Cross-Shaped, CVPR 2022项目地址: https://gitcode.com/gh_mirrors/cs/CSWin-TransformerCSWin Transformer 是微软提出的 CVPR 2022 通用视觉 Transformer 骨干网络用十字形窗口自注意力替代传统全注意力在 ImageNet-1K 上以仅 4.3G FLOPs 达到 82.8% Top-1 精度Tiny 版。本文将带你使用仓库自带的 main.py 脚本从零开始完成多卡分布式训练、EMA 权重平均、混合精度AMP的完整训练流程并逐一解读train.sh背后的关键参数。上图展示了 CSWin 的核心思想将特征图切分为十字形窗口Cross-Shaped Window水平与垂直两条条纹并行计算自注意力在极低成本下实现跨大范围的全局感知下方则是 4 阶段渐进降采样的分层架构224 输入 → 32×32 特征。一、环境安装3 步搞定依赖非常轻量核心是PyTorch 1.4timm0.3.4获取代码git clone https://gitcode.com/gh_mirrors/cs/CSWin-Transformer cd CSWin-Transformer一键安装依赖运行 install_req.sh它会安装timm0.3.4、opencv、einops、PyTurboJPEG等训练所需的全部库脚本中指定了torch 1.7.1 cu110可按自己的 CUDA 版本调整。可选安装 ApexREADME 提到混合精度微调会用到 NVIDIA Apex不过main.py同时支持原生 PyTorch AMP较新的 PyTorch 无需 Apex 即可开启--native-amp。二、准备 ImageNet 数据集数据集需按标准目录摆放且仓库已内置 1K 类别的索引文件dataset/imagenet_class_index.jsonimagenet/ ├── train/ │ ├── n01440764/ │ │ ├── n01440764_10026.JPEG │ └── ... └── val/ ├── n01440764/ │ ├── ILSVRC2012_val_00000293.JPEG └── ...训练/验证的样本清单由 dataset/ILSVRC2012_name_train.txt 和 dataset/ILSVRC2012_name_val.txt 提供labeled_memcached_dataset.py 中的McDataset会读取这两个文件并把类别名映射为 0~999 的标签——因此必须保证清单文件在仓库根目录下可访问代码里写的是相对路径./dataset/...。三、一行命令启动 8 卡分布式训练train.sh 的本质只有一行python -m torch.distributed.launch --nproc_per_node$NUM_PROC main.py $它把第一个参数当作 GPU 进程数其余参数原样转发给 main.py。以训练 CSWin-Tiny 为例bash train.sh 8 --data data path --model CSWin_64_12211_tiny_224 \ -b 256 --lr 2e-3 --weight-decay .05 --amp \ --img-size 224 --warmup-epochs 20 --model-ema-decay 0.99984 --drop-path 0.2官方为三个轻量版本提供的训练配方如下 模型模型名Batch学习率Weight DecayEMA DecayDrop-PathCSWin-TinyCSWin_64_12211_tiny_2242562e-30.050.999840.2CSWin-SmallCSWin_64_24322_small_2242562e-30.050.999840.4CSWin-BaseCSWin_96_24322_base_2241281e-30.10.999920.5显存不够时把-b降到 128、学习率减半如--lr 1e-3或追加--use-chk开启梯度检查点activation checkpointing用少量时间换大量显存。四、main.py 全流程关键步骤解读main.py借鉴了 timm/DeiT 的训练框架整个流程可以拆成 6 步分布式初始化main.py#L272-L288脚本检测环境变量WORLD_SIZE多卡时自动切换到 DDP 模式每个进程绑定一张 GPUcuda:local_rank并用 NCCL 后端init_process_group建立通信。随机种子按seed rank设置保证各卡数据不重复。模型构建main.py#L299-L317通过create_model按名字实例化模型定义在 models/cswin.py默认 1000 类、224 输入并打印参数量。混合精度AMP见下文第五节。优化器与调度器默认 AdamW 余弦退火前 20 个 epoch 线性预热到峰值学习率之后缓慢衰减到min-lr默认 1e-5。数据加载训练集启用 RandAugment、随机擦除、Mixup0.8 CutMix1.0、标签平滑0.1等主流增强且每个 epoch 会set_epoch重新打乱数据main.py#L555-L557。训练主循环每个 epoch 依次执行训练 → 原模型验证 → EMA 模型验证 → 更新学习率 → 写summary.csv→ 按 top1 保存最佳 checkpointmain.py#L554-L589。五、三大进阶机制分布式、EMA、混合精度1. 分布式训练DDP多卡模式下main.py会先处理 SyncBN可选再用原生DistributedDataParallel包装模型main.py#L397-L420注意EMA 模型不进 DDP——因为 EMA 只需在 rank 0 上维护一份。验证时各卡的 loss/accuracy 会通过reduce_tensor全局归约得到整机统一的评估指标。2. EMA模型权重的滑动平均--model-ema默认开启main.py#L196-L201。训练每一步后EMA 权重按EMA decay × EMA (1 − decay) × 当前权重缓慢更新main.py#L647-L648相当于对权重做指数平滑能显著抑制训练后期的抖动、白捡 0.2%~0.5% 精度。decay需随 batch 总量调整batch 越大、步数越少decay 越要接近 1表中 Tiny/Small 用 0.99984Base 用 0.99992。每个 epoch 结束会用EMA 模型验证并以其指标决定保存哪个 checkpoint。3. 混合精度训练AMP传--amp后脚本会优先尝试 Apex AMP检测不到则自动回落到原生 PyTorch AMPmain.py#L334-L347。原生 AMP 下由NativeScaler负责前向用autocast半精度main.py#L629-L631、损失放大、梯度反缩放后再step从而在损失几乎不变的情况下约提速 40%~60% 并减半显存占用。若环境两者都没有会打印警告并退回 float32不会报错崩溃。六、断点续训与 Checkpoint 管理续训加--resume ckpt路径即可恢复模型 优化器 学习率调度器 EMA 的完整状态main.py#L381-L386只想恢复权重、重置优化器则加--no-resume-opt。只评估--eval_checkpoint ckpt会直接加载模型跑一遍验证集并打印 Top-1 后退出main.py#L531-L535非常适合对比不同 epoch 的效果。自动保存checkpoint_saver.py 中的CheckpointSaver在 rank 0 上按验证 top1 维护model_best.*与last.*两组权重并同步写入args.yaml和summary.csv方便用 Excel 直接画训练曲线。七、训练结果速览与常见问题训练完成后不同规模在 ImageNet-1K224×224上的官方结果 模型参数量FLOPsTop-1CSWin-Tiny23M4.3G82.8CSWin-Small35M6.9G83.6CSWin-Base78M15.0G84.2新手常见疑问 FAQ❓单卡能训吗能。WORLD_SIZE不存在时自动进入单进程模式main.py#L271-L295直接python main.py --data ... --model CSWin_64_12211_tiny_224即可但吞吐会低很多。❓384 分辨率怎么训加--img-size 384官方通常采用224 预训练 → 384 微调两阶段策略微调脚本见 finetune.sh其中--finetune指定 224 预训练权重、--ema-finetune会同步微调 EMA 模型。❓日志里 EMA 和原始模型差距大训练初期属正常现象EMA 需要足够多步才跟上关注 EMA 那行日志即可最终权重以 EMA 为准。❓想迁移到语义分割仓库自带 segmentation/ 目录提供 UPerNet CSWin 的 MMSegmentation 配置segmentation/configs/cswin/upernet_cswin_base.py 等可用这里训好的权重直接当 backbone。总结用 CSWin Transformer 训练自己的分类模型只需记住一条主线train.sh拉起多进程 DDP →main.py里 AdamW 余弦学习率 混合精度 EMA 四件套 → CheckpointSaver 自动保存最优权重。对照本文第二节的配方表替换--model、-b、--lr即可复现 Tiny/Small/Base 三个规格显存紧张时优先--use-chk与降 batch。训练完成后再用finetune.sh做高分辨率精调就能拿到接近论文水平的精度了。【免费下载链接】CSWin-TransformerCSWin Transformer: A General Vision Transformer Backbone with Cross-Shaped, CVPR 2022项目地址: https://gitcode.com/gh_mirrors/cs/CSWin-Transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

相关新闻