ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ColossalAI 新 API 实战:用 Booster + Plugin 在 CIFAR-10 上从零训练 ResNet

2026/9/9 19:04:26 拓冰建站 浏览量
ColossalAI 新 API 实战:用 Booster + Plugin 在 CIFAR-10 上从零训练 ResNet 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 的新版高层训练 APIBoosterPlugin在 CIFAR-10 数据集上从零训练 ResNet-18。通过阅读本文你将掌握colossalai run的多卡启动方式、torch_ddp/torch_ddp_fp16/low_level_zero三种数据并行插件的切换方法、以及配套的 checkpoint 保存与恢复流程并能复现示例中给出的多卡训练精度。一、示例概览与目录结构该示例位于仓库 examples/tutorial/new_api/cifar_resnet 目录下与 cifar_vitViT 版、glue_bertBERT 版等共同构成 ColossalAI 新 API 的演示集合。该目录内的关键文件如下文件作用train.py训练主脚本含参数解析、插件/Booster 构造、分布式数据加载、训练循环与 checkpoint 存取eval.py单机评测脚本加载某个 epoch 的模型权重在测试集上计算 Top-1 精度test_ci.shCI 回归脚本循环用三种插件各跑一遍训练并校验目标精度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 暂不支持 ResNettorch_ddp-r, --resume从某个 epoch 的 checkpoint 恢复训练取值为整数 epoch 编号-1表示不恢复-c, --checkpointcheckpoint 保存目录./checkpoint-i, --interval每隔多少个 epoch 保存一次 checkpoint设为0表示不保存5--target_acc目标精度训练结束时若未达到该精度则抛出 AssertionErrorNone不校验评测参数eval.py参数说明默认值-e, --epoch指定加载哪个 epoch 的模型权重对应model_{epoch}.pth80-c, --checkpointcheckpoint 所在目录./checkpointeval.py中的模型同样使用torchvision.models.resnet18(num_classes10)并在加载.cuda()后执行注意评测脚本需要读取{checkpoint}/model_{epoch}.pth这一权重文件因此应使用与训练一致的 checkpoint 目录与 epoch 编号。三、环境安装与数据准备安装依赖pip install -r requirements.txtrequirements.txt 仅包含四个包colossalai、torch、torchvision、tqdm。CIFAR-10 数据集无需手动下载训练脚本会通过torchvision.datasets.CIFAR10自动完成下载。数据集路径数据集根目录可通过环境变量DATA指定见 train.pydata_path os.environ.get(DATA, ./data)即在 test_ci.sh 中设置为export DATA/data/scratch/cifar-10若未设置DATA则默认落在当前目录下的./data。数据下载由 train.py 中coordinator.priority_execution()保护——该上下文保证下载动作只由优先级较高的进程执行避免多进程并发写同一目录产生冲突。训练侧使用了经典的数据增强流水线见 train.pytransform_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 原生DistributedDataParallelfp16 混合精度 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-3Batch size100见build_dataloader(100, coordinator, plugin)调用优化器HybridAdam即 colossalai.nn.optimizer 提供的混合 Adam学习率调度MultiStepLR(optimizer, milestones[20, 40, 60, 80], gamma1/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 中给出的多卡训练精度参考如下ModelSingle-GPU Baseline FP32Booster DDP FP32Booster DDP FP16Booster Low Level ZeroResNet-1885.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 并构造 Boosterbooster_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_policystatic, strict_ddp_modeTrue, initial_scale2**5) elif args.plugin low_level_zero: plugin LowLevelZeroPlugin(initial_scale2**5) booster Booster(pluginplugin, **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_scale2**5是混合精度动态 loss scaling 的初始值。需要留意两种 fp16 相关插件torch_ddp_fp16、low_level_zero都会进行混合精度训练其中 low_level_zero 在LowLevelZeroPlugin(initial_scale2**5)内部同时启用了 fp16且并未在booster_kwargs里再传mixed_precision。3. 构造分布式 DataLoadertrain_dataloader plugin.prepare_dataloader(train_dataset, batch_sizebatch_size, shuffleTrue, drop_lastTrue) test_dataloader plugin.prepare_dataloader(test_dataset, batch_sizebatch_size, shuffleFalse, drop_lastFalse)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, criterioncriterion, lr_schedulerlr_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进度条同样通过disablenot coordinator.is_master()只在主进程显示。七、插件切换背后的设计思想这个示例最大价值在于演示了 ColossalAI 新 API 的一行切换并行方案能力。从实现上看Plugin负责并行策略DDP、ZeRO 等、数据加载、模型与优化器包装、checkpoint I/O 与 LoRA/无同步等高级能力例如TorchDDPPlugin还支持fp8_communicationFP8 梯度通信压缩 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_classes10)换成自己的模型注意分类头维度与数据集类别数一致将 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),仅供参考