ColossalAI 新 API 实战:用 Booster + Plugin 在 CIFAR-10 上从零训练 ResNet
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
导读
本教程基于 ColossalAI 仓库中的 examples/tutorial/new_api/cifar_resnet/README.md 实战示例,讲解如何利用 ColossalAI 的新版高层训练 API(Booster+Plugin)在 CIFAR-10 数据集上从零训练 ResNet-18。通过阅读本文,你将掌握colossalai run的多卡启动方式、torch_ddp/torch_ddp_fp16/low_level_zero三种数据并行插件的切换方法、以及配套的 checkpoint 保存与恢复流程,并能复现示例中给出的多卡训练精度。
一、示例概览与目录结构
该示例位于仓库 examples/tutorial/new_api/cifar_resnet 目录下,与 cifar_vit(ViT 版)、glue_bert(BERT 版)等共同构成 ColossalAI 新 API 的演示集合。该目录内的关键文件如下:
| 文件 | 作用 |
|---|---|
train.py | 训练主脚本,含参数解析、插件/Booster 构造、分布式数据加载、训练循环与 checkpoint 存取 |
eval.py | 单机评测脚本,加载某个 epoch 的模型权重在测试集上计算 Top-1 精度 |
test_ci.sh | CI 回归脚本,循环用三种插件各跑一遍训练并校验目标精度 |
requirements.txt | 运行依赖:colossalai、torch、torchvision、tqdm |
需要说明的是,该目录属于新 API 教程范畴。仓库 examples/tutorial/new_api/README.md 明确指出该 API 仍处于密集开发中,尚未正式发布,因此阅读与复现时需留意仓库内 API 的演进可能造成差异。
二、命令行参数说明
train.py与eval.py使用argparse解析参数,参数在 train.py 与 eval.py 中声明。
训练参数(train.py)
| 参数 | 说明 | 默认值 |
|---|---|---|
-p, --plugin | 使用的数据并行插件,可选torch_ddp、torch_ddp_fp16、low_level_zero(代码中的 choices 还预留了gemini,但注释标注 gemini 暂不支持 ResNet) | torch_ddp |
-r, --resume | 从某个 epoch 的 checkpoint 恢复训练,取值为整数 epoch 编号 | -1,表示不恢复 |
-c, --checkpoint | checkpoint 保存目录 | ./checkpoint |
-i, --interval | 每隔多少个 epoch 保存一次 checkpoint;设为0表示不保存 | 5 |
--target_acc | 目标精度,训练结束时若未达到该精度则抛出 AssertionError | None(不校验) |
评测参数(eval.py)
| 参数 | 说明 | 默认值 |
|---|---|---|
-e, --epoch | 指定加载哪个 epoch 的模型权重(对应model_{epoch}.pth) | 80 |
-c, --checkpoint | checkpoint 所在目录 | ./checkpoint |
eval.py中的模型同样使用torchvision.models.resnet18(num_classes=10)并在加载.cuda()后执行,注意评测脚本需要读取{checkpoint}/model_{epoch}.pth这一权重文件,因此应使用与训练一致的 checkpoint 目录与 epoch 编号。
三、环境安装与数据准备
安装依赖
pip install -r requirements.txtrequirements.txt 仅包含四个包:colossalai、torch、torchvision、tqdm。CIFAR-10 数据集无需手动下载,训练脚本会通过torchvision.datasets.CIFAR10自动完成下载。
数据集路径
数据集根目录可通过环境变量DATA指定,见 train.py:
data_path = os.environ.get("DATA", "./data")即在 test_ci.sh 中设置为export DATA=/data/scratch/cifar-10;若未设置DATA,则默认落在当前目录下的./data。数据下载由 train.py 中coordinator.priority_execution()保护——该上下文保证下载动作只由优先级较高的进程执行,避免多进程并发写同一目录产生冲突。
训练侧使用了经典的数据增强流水线,见 train.py:
transform_train = transforms.Compose( [transforms.Pad(4), transforms.RandomHorizontalFlip(), transforms.RandomCrop(32), transforms.ToTensor()] ) transform_test = transforms.ToTensor()即 4 像素填充 + 随机水平翻转 + 随机裁剪(裁剪回 32×32)+ 转 Tensor,测试集不做增强。
四、快速开始:三种插件的训练命令
在安装好依赖后,直接使用colossalai run启动多进程分布式训练即可。
训练
# train with torch DDP with fp32 colossalai run --nproc_per_node 2 train.py -c ./ckpt-fp32 # train with torch DDP with mixed precision training colossalai run --nproc_per_node 2 train.py -c ./ckpt-fp16 -p torch_ddp_fp16 # train with low level zero colossalai run --nproc_per_node 2 train.py -c ./ckpt-low_level_zero -p low_level_zero三条命令分别对应三种训练配置,--nproc_per_node 2表示单机 2 卡(多机扩展方式可参考 cli/launcher 的 runner 实现):
- fp32 全精度 DDP:使用默认插件
torch_ddp,对应 PyTorch 原生DistributedDataParallel; - fp16 混合精度 DDP:
-p torch_ddp_fp16,在 DDP 之上叠加 FP16 混合精度; - 低阶 ZeRO:
-p low_level_zero,即 ZeRO-1/2 风格的分片优化(Low Level Zero)。
CI 脚本 test_ci.sh 展示了同样的三种插件组合(用 4 卡、--interval 0关闭存盘、--target_acc 0.84校验精度 ≥84%):
for plugin in "torch_ddp" "torch_ddp_fp16" "low_level_zero"; do colossalai run --nproc_per_node 4 train.py --interval 0 --target_acc 0.84 --plugin $plugin done该脚本为理解"如何把训练接进自动回归"提供了直接范例。
评测
# evaluate fp32 training python eval.py -c ./ckpt-fp32 -e 80 # evaluate fp16 mixed precision training python eval.py -c ./ckpt-fp16 -e 80 # evaluate low level zero training python eval.py -c ./ckpt-low_level_zero -e 80每个 checkpoint 目录下保存了model_{epoch}.pth,-e 80表示评测第 80 个 epoch 结束时保存的权重。
五、训练超参数与预期精度
核心超参数
训练超参数硬编码在 train.py:
- 总训练轮数
NUM_EPOCHS = 80; - 基础学习率
LEARNING_RATE = 1e-3; - Batch size:100(见
build_dataloader(100, coordinator, plugin)调用); - 优化器:
HybridAdam,即 colossalai.nn.optimizer 提供的混合 Adam; - 学习率调度:
MultiStepLR(optimizer, milestones=[20, 40, 60, 80], gamma=1/3)。
值得注意的细节是线性学习率缩放。在分布式环境初始化后,脚本做了如下处理(train.py):
# 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即学习率随参与训练的 GPU 总数线性放大,以在增大 batch 的同时保持收敛行为,这与大批量训练常用的线性缩放规则(Linear Scaling Rule)一致。
预期精度
README 中给出的多卡训练精度参考如下:
| Model | Single-GPU Baseline FP32 | Booster DDP FP32 | Booster DDP FP16 | Booster Low Level Zero |
|---|---|---|---|---|
| ResNet-18 | 85.85% | 84.91% | 85.46% | 84.50% |
其中单卡基线改编自 pytorch-tutorial 的 ResNet-CIFAR-10 脚本,并将网络替换为torchvision.models.resnet18。需要注意:该表是示例作者在特定软硬件环境下测得的结果,仅用于横向对比三种插件在精度上的等价性(三者互有细微高低,均处于正常范围),不应理解为绝对的性能承诺;你在自己环境中的实际数值会因随机种子、设备与软件版本而浮动。
六、深入源码:train.py 的训练流程拆解
下面按 train.py 的执行顺序,拆解新 API 的核心调用链,帮助你理解"一段普通 PyTorch 训练代码是如何被改造成分布式可扩展训练的"。
1. 启动分布式环境
colossalai.launch_from_torch() coordinator = DistCoordinator()colossalai.launch_from_torch()从torch.distributed已初始化的环境(由colossalai run或torchrun建立)中获取 rank/world_size 等信息完成初始化;随后DistCoordinator(见 colossalai/cluster/dist_coordinator.py)封装了"当前进程是否为 master"、"世界规模多大"等常用查询,例如coordinator.is_master()控制日志打印、coordinator.priority_execution()控制数据下载等单次任务。
2. 选择 Plugin 并构造 Booster
booster_kwargs = {} if args.plugin == "torch_ddp_fp16": booster_kwargs["mixed_precision"] = "fp16" if args.plugin.startswith("torch_ddp"): plugin = TorchDDPPlugin() elif args.plugin == "gemini": plugin = GeminiPlugin(placement_policy="static", strict_ddp_mode=True, initial_scale=2**5) elif args.plugin == "low_level_zero": plugin = LowLevelZeroPlugin(initial_scale=2**5) booster = Booster(plugin=plugin, **booster_kwargs)从 colossalai/booster/plugin/torch_ddp_plugin.py 的类定义可见,TorchDDPPlugin本质是对 PyTorchDistributedDataParallel的封装:在configure()中先将模型搬到当前设备并转换SyncBatchNorm,再用TorchDDPModel包裹模型。插件与 Booster 的设计将"并行方案""混合精度""checkpoint I/O"等横切关注点解耦,用户只需替换 plugin 与精度参数,训练循环几乎无需改动。
LowLevelZeroPlugin则对应 ZeRO 的 low-level 实现(见 colossalai/booster/plugin/low_level_zero_plugin.py),其构造入参中initial_scale=2**5是混合精度动态 loss scaling 的初始值。需要留意:两种 fp16 相关插件(torch_ddp_fp16、low_level_zero)都会进行混合精度训练,其中 low_level_zero 在LowLevelZeroPlugin(initial_scale=2**5)内部同时启用了 fp16,且并未在booster_kwargs里再传mixed_precision。
3. 构造分布式 DataLoader
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)prepare_dataloader定义在基类 colossalai/booster/plugin/dp_plugin_base.py:它依据当前world_size与rank为每个进程自动装配DistributedSampler,从而保证每张卡看到互不重叠的数据分片,并提供可复现的seed_worker。也就是说,我们无需手写数据切分逻辑,插件已经替我们完成。
4. Boost 模型、优化器与调度器
model, optimizer, criterion, _, lr_scheduler = booster.boost( model, optimizer, criterion=criterion, lr_scheduler=lr_scheduler )Booster.boost(见 colossalai/booster/booster.py)是整套 API 的中枢:它会调用 plugin 的configure()对模型进行并行化改造、根据mixed_precision配置精度、并返回经过包装的 optimizer/lr_scheduler 等对象。之后训练循环中应使用返回的对象。
5. Checkpoint 存取与恢复
恢复与保存统一使用 Booster 暴露的接口:
# 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") # save (每隔 interval 个 epoch) 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")以TorchDDPPlugin为例,其配套的TorchDDPCheckpointIO(同文件内定义)重写了模型/优化器/scheduler 的存与取,并将真正的落盘限定在 master 进程上执行(保存前判断coordinator.is_master()),避免多进程重复写盘造成竞争。恢复训练时,start_epoch从args.resume开始续跑:
start_epoch = args.resume if args.resume >= 0 else 0 for epoch in range(start_epoch, NUM_EPOCHS): ...由此实现"断电续训"能力——例如中断在第 60 epoch,可用-r 60从model_60.pth、optimizer_60.pth、lr_scheduler_60.pth恢复。
6. 反向传播入口
训练循环中前向计算与普通 PyTorch 完全一致,唯一的关键差异是反向传播:
booster.backward(loss, optimizer) optimizer.step() optimizer.zero_grad()Booster.backward会按当前插件与精度配置,正确执行梯度缩放/规约等操作(混合精度场景下对应 GradScaler 的scale逻辑),随后仍是标准的optimizer.step()与zero_grad()。
7. 分布式精度统计
评测函数evaluate展示了一个在多卡环境下正确统计精度的通用写法(train.py):每张卡各自统计correct与total张量,再通过dist.all_reduce汇总到全体进程,最后仅由 master 进程打印,避免多进程重复输出造成日志混乱。train_epoch中的tqdm进度条同样通过disable=not coordinator.is_master()只在主进程显示。
七、插件切换背后的设计思想
这个示例最大价值在于演示了 ColossalAI 新 API 的"一行切换并行方案"能力。从实现上看:
Plugin负责并行策略(DDP、ZeRO 等)、数据加载、模型与优化器包装、checkpoint I/O 与 LoRA/无同步等高级能力;例如TorchDDPPlugin还支持fp8_communication(FP8 梯度通信压缩 hook,见其构造函数)。Booster负责编排 plugin 与 mixed precision,向用户暴露统一且稳定的boost/backward/save_*/load_*接口。
因此用户可以在"训练代码几乎不动"的前提下,从torch_ddp平移到torch_ddp_fp16获得显存/吞吐收益,或切换到low_level_zero以支持更大模型的低阶 ZeRO 分片训练。对于需要更高阶能力的场景(Gemini 显存卸载、混合并行、流水线并行等),仓库 colossalai/booster/plugin 下还提供了GeminiPlugin、HybridParallelPlugin、MoeHybridParallelPlugin、TorchFSDPPlugin等更多插件,均遵循同一套Plugin基类契约,可参考 examples/tutorial/new_api 下其他示例与对应插件源码继续探索。
八、从示例迁移到自己的训练任务
若要将本示例改写成自己的训练脚本,只需替换以下"业务相关"部分,分布式样板代码可整体保留:
- 将
torchvision.models.resnet18(num_classes=10)换成自己的模型,注意分类头维度与数据集类别数一致; - 将 CIFAR-10 的数据加载与增强替换为自己的 Dataset/Transform;
- 保留
colossalai.launch_from_torch()→ 构造 plugin →Booster(...)→plugin.prepare_dataloader(...)→booster.boost(...)→booster.backward(...)的骨架; - 多卡线性学习率缩放、master-only 日志打印、
dist.all_reduce汇总指标、按 interval 用booster.save_*存盘并用-r恢复等实践,均可按需沿用。
九、常见问题与注意事项
- 数据集下载并发冲突:务必保留
coordinator.priority_execution()包裹下载逻辑,否则多进程可能同时写数据集目录;或者预先用DATA环境变量指向已下载好的数据。 - eval 与 train 的 epoch 对应关系:
eval.py -e 80加载的是model_80.pth;若训练因中断未跑满 80 epoch,或--interval非 5,请按实际存在的 checkpoint 编号评测。 - checkpoint 命名与目录:训练脚本会在
--interval > 0时创建目录(train.py);若--interval 0,整个训练过程不会产生任何权重文件,后续无法评测。 - gemini 插件的使用限制:train.py 的代码分支虽然保留了
gemini选项,但注释明确标注 "gemini is not supported resnet now",当前示例不建议选用。 - API 处于演进期:本示例位于新 API 演示目录,
Booster/Plugin相关接口可能随版本调整,遇到差异时以当前仓库 colossalai/booster 下的源码实现为准。
总而言之,通过本示例你可以完整走通一条"用 ColossalAI 新 API 从零训一个 CNN 分类器"的路径:从colossalai run拉起多卡,到以三种数据并行插件快速横向对比,再到 checkpoint 的保存、恢复与单卡评测。这套以 Booster 为中心的代码骨架,也正是后续学习 Gemini、混合并行乃至大模型预训练等进阶能力的基础。
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考