news 2026/9/10 15:58:12

ColossalAI 分布式训练 ResNet-18 于 CIFAR-10:从 torch DDP、混合精度到 Low Level ZeRO 的完整实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
ColossalAI 分布式训练 ResNet-18 于 CIFAR-10:从 torch DDP、混合精度到 Low Level ZeRO 的完整实战

ColossalAI 分布式训练 ResNet-18 于 CIFAR-10:从 torch DDP、混合精度到 Low Level ZeRO 的完整实战

【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI

本指南基于 ColossalAI 官方示例 examples/images/resnet 展开,讲解如何使用一套几乎无需改动训练逻辑的脚本,通过 ColossalAI Booster 在 CIFAR-10 上从头训练 ResNet-18,并一键切换torch_ddptorch_ddp_fp16low_level_zerogemini四种分布式/优化插件。读完本文,你将掌握示例提供的训练参数与断点恢复机制、运行与评估命令、各插件背后的源码调用关系,以及如何在多卡环境下保持模型精度的完整实验方法。

示例概览与文件结构

该示例位于仓库 examples/images/resnet 目录,共包含五个文件:

  • train.py:主训练脚本,负责分布式环境初始化、数据集构建、模型/优化器创建,并借助 Booster 完成多插件训练。
  • eval.py:独立的单机评估脚本,加载某个 epoch 保存的模型权重并计算 CIFAR-10 测试集准确率。
  • requirements.txt:运行依赖清单,包含colossalaitorchtorchvisiontqdmpytest
  • test_ci.sh:CI 冒烟测试脚本模板,内部以注释形式给出了“不同插件 × 目标精度 0.84”的验证意图。
  • README.md:本文所述的官方使用说明。

示例的训练逻辑对 torchvision 用户而言非常直观:数据集采用torchvision.datasets.CIFAR10,模型采用torchvision.models.resnet18(num_classes=10),损失函数为交叉熵,优化器则是 ColossalAI 封装的HybridAdam。分布式训练能力并非通过改写模型实现,而是由 Booster 这一统一入口在启动阶段注入,这正是理解该示例的关键。

环境准备与依赖安装

在运行任何训练/评估脚本前,需要先完成依赖安装:

pip install -r examples/images/resnet/requirements.txt

需求文件内容即:

colossalai torch torchvision tqdm pytest

其中colossalai是本仓库主包(对应仓库根目录的 setup.py),安装时应从仓库根目录执行pip install .或直接使用已安装的发行版本;pytest主要用于 CI 场景而非训练本身。

CIFAR-10 数据集无需手动下载:训练脚本会在首次运行时自动下载,下载根目录可通过环境变量DATA指定,默认落在./data(详见 train.py 的data_path = os.environ.get("DATA", "./data"))。CI 脚本 test_ci.sh 中即通过export DATA=/data/scratch/cifar-10来指定数据集缓存位置。

训练脚本命令行参数详解

示例同时提供训练与评估两套参数体系,这里结合 train.py 与 eval.py 的 argparse 定义逐项说明。

训练参数(train.py)

参数全称类型默认值说明
-p--pluginstrtorch_ddp使用的分布式/优化插件,可选值为torch_ddptorch_ddp_fp16low_level_zerogemini
-r--resumeint-1从第几个 epoch 的断点恢复训练,-1表示不恢复
-c--checkpointstr./checkpoint保存 checkpoint 的目录
-i--intervalint5每隔多少个 epoch 保存一次 checkpoint;设为0则完全不保存
--target_acc-floatNone目标测试精度,训练结束时若未达到则抛异常(供 CI 回归使用)

几点值得注意的细节:

  • -p的可选值来源于 train.py 的choices列表,README 表格只列出前三种,但源码层面gemini亦在合法范围内,并带有一条FIXME(ver217): gemini is not supported resnet now的注释,说明 Gemini 路径仍处于实验性阶段。
  • -r的取值直接对应 checkpoint 文件名的 epoch 编号。例如-r 40会加载model_40.pthoptimizer_40.pthlr_scheduler_40.pth三个文件(见下文“断点保存与恢复”小节)。
  • --target_acc配合脚本末尾的assert accuracy >= args.target_acc(train.py)使用,是 CI 自动化验证精度的关键开关。

评估参数(eval.py)

参数全称默认值说明
-e--epoch80选择评估哪个 epoch 保存的模型权重
-c--checkpoint./checkpointcheckpoint 所在目录

评估脚本为单进程单卡设计:它加载{checkpoint}/model_{epoch}.pth到 CUDA 设备,以 batch size 128 遍历测试集,最终打印形如Accuracy of the model on the test images: xx.xx %的结果(eval.py)。因此,训练阶段为分布式保存的模型权重(每卡保存全量模型权重),可在训练完成后由任意单机脚本独立评估。

三步完成分布式训练

官方推荐通过 ColossalAI 自带的colossalai run启动器拉起多进程。目录不存在时会由脚本自动创建(Path(args.checkpoint).mkdir(parents=True, exist_ok=True))。

1. 以 torch DDP + FP32 训练

colossalai run --nproc_per_node 2 train.py -c ./ckpt-fp32

默认插件即torch_ddp,对应代码中TorchDDPPlugin()的实例化,仅做数据并行包装,保持 FP32 精度基线。学习率会随进程数线性放大,见下文源码解读。

2. 以 torch DDP + FP16 混合精度训练

colossalai run --nproc_per_node 2 train.py -c ./ckpt-fp16 -p torch_ddp_fp16

选择该插件时,train.py 会向 Booster 传入mixed_precision="fp16",由 Booster 内部的混合精度工具(对应仓库 colossalai/booster/mixed_precision 子模块)接管前向/反向与梯度缩放,从而在几乎不损失精度的前提下显著降低显存与带宽压力。

3. 以 Low Level ZeRO 训练

colossalai run --nproc_per_node 2 train.py -c ./ckpt-low_level_zero -p low_level_zero

该路径实例化LowLevelZeroPlugin(initial_scale=2**5)(train.py)。initial_scale=2**5是 FP16 动态损失缩放因子的初值;此插件对应仓库 colossalai/booster/plugin/low_level_zero_plugin.py 的实现,在不修改模型定义的前提下完成优化器状态的分片与通信,是比纯 DDP 更省显存的替代方案。

启动器与运行时初始化链路

上述命令最终都会让 train.py 执行colossalai.launch_from_torch()DistCoordinator()launch_from_torch定义于 colossalai/initialize.py,它从 PyTorch 启动器写入的环境变量(RANKLOCAL_RANKWORLD_SIZEMASTER_ADDRMASTER_PORT)中读取进程拓扑信息并完成通信后端初始化,因而colossalai run与标准的torchrun语义保持一致。

训练完成后执行评估

模型经过 80 个 epoch 训练后,每个-i指定间隔都会留下权重快照。针对上述三套训练分别执行:

# 评估 FP32 训练结果 python eval.py -c ./ckpt-fp32 -e 80 # 评估 FP16 混合精度训练结果 python eval.py -c ./ckpt-fp16 -e 80 # 评估 low level zero 训练结果 python eval.py -c ./ckpt-low_level_zero -e 80

注意评估脚本默认参数-e 80正好对应用满 80 个 epoch 的最终权重;若只训练了 40 个 epoch 或希望评测中间快照,把-e改成对应 epoch 编号即可。

预期精度表现与基线说明

README 给出了在多卡训练下可复现的精度参考值(以 ResNet-18 为模型):

ModelSingle-GPU Baseline FP32Booster DDP FP32Booster DDP FP16Booster Low Level ZeroBooster Gemini
ResNet-1885.85%84.91%85.46%84.50%84.60%

几点事实澄清与注意:

  • 单卡 FP32 基线 85.85% 是 README 声明值,其来源为将 PyTorch 官方教程《CNN ResNet for CIFAR-10》脚本改造为使用torchvision.models.resnet18后的结果(README 底部 Note 明确注明此出处)。
  • 三种 Booster 方案(FP32 DDP、FP16 DDP、Low Level Zero)的测试集精度均落在 84%~86% 区间,与单卡基线差异在 1 个百分点以内,说明分布式/混合精度/ZeRO 优化并不会显著牺牲模型收敛质量。
  • 表格中 Gemini 一行对应精度 84.60%,但正如前文所述,train.py 中存在 “gemini is not supported resnet now” 的 FIXME 注释——该插件路径当前视为实验性支持,复现时请以实际运行输出为准,切勿将表格数值当作绝对承诺。

源码级纵深:脚本内部的关键设计

为帮助读者真正理解这套“零侵入式”分布式训练是如何做到的,这里沿 train.py 的执行顺序拆解五个内部要点。

1. 学习率的线性缩放

# update the learning rate with linear scaling # old_gpu_num / old_lr = new_gpu_num / new_lr global LEARNING_RATE LEARNING_RATE *= coordinator.world_size

基线学习率 1e-3 按 GPU 数量线性放大(train.py),这是多卡同步 SGD 场景下保证“等效批次大小不变、收敛行为不变”的常用经验法则;coordinator.world_size来自DistCoordinator,其底层即torch.distributed.get_world_size()

2. 数据加载由插件接管

train_dataloader = plugin.prepare_dataloader(train_dataset, batch_size=batch_size, shuffle=True, drop_last=True) test_dataloader = plugin.prepare_dataloader(test_dataset, batch_size=batch_size, shuffle=False, drop_last=False)

build_dataloader中训练集使用Pad(4) + RandomHorizontalFlip + RandomCrop(32)增强(batch size 100),测试集仅做ToTensor()。分布式的数据分片(shuffle/sampler)被封装进各插件的prepare_dataloader,因此业务代码无需自行构造DistributedSampler。数据集下载被包在coordinator.priority_execution()上下文里,保证多进程同时就绪、仅主进程执行下载等易冲突操作。

3. 插件选择与 Booster 组装

if args.plugin.startswith("torch_ddp"): plugin = TorchDDPPlugin() elif args.plugin == "gemini": plugin = GeminiPlugin(initial_scale=2**5) elif args.plugin == "low_level_zero": plugin = LowLevelZeroPlugin(initial_scale=2**5) booster = Booster(plugin=plugin, **booster_kwargs)

在 torch DDP 分支中,FP16 由mixed_precision="fp16"这个独立维度开启,而torch_ddptorch_ddp_fp16共用同一TorchDDPPlugin——这种“并行策略插件 × 混合精度开关”正交组合的设计贯穿整个 ColossalAI 新 Booster API。

4. 优化器与学习率调度

  • 优化器为HybridAdam(从 colossalai.nn.optimizer 导入),支持与 ZeRO/Gemini 的分布式优化器状态协作。
  • 调度器为MultiStepLR(optimizer, milestones=[20, 40, 60, 80], gamma=1/3)(train.py),即在第 20/40/60/80 epoch 学习率降至原来的 1/3,这也是 CIFAR-10 图像分类任务中的经典阶梯式衰减配置。
  • booster.boost(model, optimizer, criterion=criterion, lr_scheduler=lr_scheduler)统一返回被“增强”后的四件套,后续代码对返回值的使用方式与普通 PyTorch 训练完全一致。

5. 断点保存与恢复

断点保存与恢复完全围绕 Booster 的四个接口展开:

# 恢复(epoch == args.resume) booster.load_model(model, f"{args.checkpoint}/model_{args.resume}.pth") booster.load_optimizer(optimizer, f"{args.checkpoint}/optimizer_{args.resume}.pth") booster.load_lr_scheduler(lr_scheduler, f"{args.checkpoint}/lr_scheduler_{args.resume}.pth") # 保存((epoch+1) % args.interval == 0) booster.save_model(model, f"{args.checkpoint}/model_{epoch + 1}.pth") booster.save_optimizer(optimizer, f"{args.checkpoint}/optimizer_{epoch + 1}.pth") booster.save_lr_scheduler(lr_scheduler, f"{args.checkpoint}/lr_scheduler_{epoch + 1}.pth")

恢复成功后start_epoch = args.resume,训练循环从断点 epoch 无缝继续。该示例分开保存模型/优化器/调度器三份状态而非打包为一个文件——正因如此,独立的 eval.py 只需读取model_{epoch}.pth即可完成单机评估。模型与优化器状态的分布式保存/加载逻辑由 colossalai/checkpoint_io 子模块承载。

与仓库内姊妹示例的横向参照

在 examples/tutorial/new_api/cifar_resnet 下存在一份内容几乎相同的教程示例(README 与 train.py 高度同源,仅精度表少列 Gemini、命令行缺少gemini选项)。若希望横向对比不同写法或查阅另一份文档化说明,可前往该目录阅读。两者共同展示了 ColossalAI 示例库中“图像模型 + CIFAR 级数据 + 数据并行/ZeRO”的标准样板,可视为同一套 API 在不同示例目录中的复现。

小结

从本示例可以提炼出可直接复用到其他 CV 任务的三步法:其一,用colossalai.launch_from_torch() + DistCoordinator完成进程初始化;其二,仅靠-p参数在torch_ddp/torch_ddp_fp16/low_level_zero/gemini间切换,而无需改动任何模型与训练循环代码;其三,借助-i周期性保存的 checkpoint,用 eval.py 在任意时刻对某一 epoch 的权重做独立评估。在 2 卡环境下,上述三种主流方案的 CIFAR-10 测试精度均能稳定复现在 84%~86% 区间,验证了 ColossalAI Booster 在分布式扩展与精度保持之间取得了良好平衡。

复现提示:精度表中的数值基于 README 给定的 ResNet-18 与 CIFAR-10 默认配置,实际结果会受随机种子、GPU 数量(影响线性缩放后的学习率)与框架版本影响;如需自动化回归,可将 CI 中用到的--target_acc 0.84附加到训练命令之后作为精度门槛。

【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 15:56:57

51单片机步进电机阀门控制器设计

简介:本资源是一套基于STC89C52单片机的电动阀门步进电机控制器完整设计资料,面向嵌入式初学者、课程设计学生及自动化控制实践者,解决工业阀门精准驱动与人机交互控制的实际问题。包内共23个文件,涵盖KEIL4源代码(.c/…

作者头像 李华
网站建设 2026/9/10 15:54:53

OpenGL核心模式实现3D小桌模型渲染

1. 项目概述:用OpenGL构建3D小桌模型在计算机图形学领域,OpenGL一直是桌面端3D渲染的工业标准。最近我在重构一个老旧的3D建模工具时,决定用纯C和OpenGL核心模式重新实现基础渲染管线,其中测试案例选择了看似简单却包含多种图形学…

作者头像 李华
网站建设 2026/9/10 15:54:20

Android康复训练APP开发:自闭症干预技术实践

1. 项目背景与核心价值自闭症谱系障碍(ASD)是一种神经发育性疾病,全球每54名儿童中就有1名患者。传统康复训练存在资源分布不均、训练成本高、家长参与度低等痛点。这个基于Android的康复训练APP正是为了解决以下核心问题:训练场景…

作者头像 李华
网站建设 2026/9/10 15:53:48

2026年网站设计开发公司怎么选?六大服务形态盘点

1. 为什么2026年还要认真选一家网站建设公司1.1 建站门槛降低了,专业门槛反而更高先说一个反直觉的现状。十年前,一个没有技术团队的企业要做官网,基本只能找外包公司报价,一套企业站动辄两三万,贵也贵得有道理&#x…

作者头像 李华