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

发布时间:2026/9/9 19:04:27

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),仅供参考

相关新闻

8G老电脑多开网页也不卡?开源内存清理工具原理与实操指南

8G老电脑多开网页也不卡?开源内存清理工具原理与实操指南

2026/9/9 19:04:27

作为一个常年和“电脑越用越卡”作斗争的人,我对标题里“8G 老电脑多开十个网页也不卡”这句话太有共鸣了。过去几年,我在各平台刷到过不少类似标题的项目,核心指的都是 GitHub 上那批轻量级开源内存清理工具。今天就专门写一篇长文&#xff…

如何在 AI 客户端中接入 Netdata Cloud 的 MCP Server?

如何在 AI 客户端中接入 Netdata Cloud 的 MCP Server?

2026/9/9 18:54:26

如何在 AI 客户端中接入 Netdata Cloud 的 MCP Server? 【免费下载链接】netdata The fastest path to AI-powered full stack observability, even for lean teams. 项目地址: https://gitcode.com/GitHub_Trending/ne/netdata 如果你的 Claude Code、Curso…

AI 智能体长期记忆实现:基于向量检索与对话存储的完整方案

AI 智能体长期记忆实现:基于向量检索与对话存储的完整方案

2026/9/9 18:54:26

在实际部署 AI 智能体时,一个最常见的抱怨是:它就像金鱼一样,只有七秒记忆。Hermes 智能体可以处理复杂的多轮对话,但一旦会话结束,它就不会记住用户之前的偏好、兴趣和结论。要让 AI 越用越聪明,不能只靠更…

SciPy科学计算环境搭建与核心模块实战:从pip安装到积分优化插值

SciPy科学计算环境搭建与核心模块实战:从pip安装到积分优化插值

2026/9/9 19:44:28

很多人第一次看到“1.5、1.7、1.13 scipy”这种标题,第一反应大概率是懵的。可能是某本教程的章节编号,也可能是SciPy三个不同子模块在学习路径上的记号。但不管它具体指什么,真正落到实际操作上,绕不开两件事:该怎么装…

深入解析燃料电池ECMS能量管理策略:从原理到工程落地

深入解析燃料电池ECMS能量管理策略:从原理到工程落地

2026/9/9 19:44:28

1. 为什么偏偏是ECMS:燃料电池能量管理的选型思路 1.1 能量管理到底在管什么 很多刚接触燃料电池系统的朋友,第一反应是“燃料电池不就是发电的吗,直接把电送到电机不就行了”。真做起来就会发现,事情远没有那么简单。燃料电池电…

鸿蒙原生应用 HarmonyOS 6.0 轻量布局:问卷首页的横向 Chip 与状态问卷卡

鸿蒙原生应用 HarmonyOS 6.0 轻量布局:问卷首页的横向 Chip 与状态问卷卡

2026/9/9 19:44:28

鸿蒙原生应用 HarmonyOS 6.0 轻量布局:问卷首页的横向 Chip 与状态问卷卡 App 21「在线问卷调查」首页(HomeTab),主题色 #3B82F6 蓝色,4 个 Tab 分别为首页(📊)、创建(✏…

Windows 10 Build 9916虚拟机安装指南:解决VMware崩溃与蓝屏问题

Windows 10 Build 9916虚拟机安装指南:解决VMware崩溃与蓝屏问题

2026/9/9 19:44:28

这次我们来看一个非常冷门的内测系统:Windows 10 Build 9916。它属于 Windows 10 早期技术预览阶段的一个内部构建版本,既没有进入正式 Windows 10 发布体系,也不是常规年度更新版本。很多人第一次拿到它,是想在虚拟机上看看 Wind…

Windows 10 Build 9916虚拟机安装崩溃排查与VMware参数配置指南

Windows 10 Build 9916虚拟机安装崩溃排查与VMware参数配置指南

2026/9/9 19:44:28

这次我们来看一个冷门内测系统——Windows 10 Build 9916。从 Build 号看,它属于 Windows 10 正式版发布前的早期技术预览版本,流通量不大,稳定性也不如后来的公开预览版。很多人在物理机里装它都会翻车,于是想到虚拟机里试&#…

Python学习路线之第二阶段:从基础语法到写出实用脚本

Python学习路线之第二阶段:从基础语法到写出实用脚本

2026/9/9 19:34:28

学Python最尴尬的瞬间,不是报错,而是安装完成之后盯着屏幕不知道下一步该干什么。下载、装环境、跑了个hello world,然后呢?刷了两天教程,代码跟着敲了一遍,合上电脑又什么都不记得。这不是你的问题&#x…

中国人民大学杨琳团队《Nature Communications》 | 全球潮汐湿地土壤有机碳时空格局与环境驱动:一项2009-2020年的全球评估

中国人民大学杨琳团队《Nature Communications》 | 全球潮汐湿地土壤有机碳时空格局与环境驱动:一项2009-2020年的全球评估

2026/9/9 1:14:29

本文首发于“生态学者”!从“湿地面积”到“土壤碳密度”:为什么需要重新认识潮汐湿地蓝碳变化?潮汐湿地位于陆地与海洋的交汇地带,包括红树林、盐沼和潮滩,是全球重要的蓝碳生态系统。其土壤能够长期储存大量有机碳&a…

adb抓包

adb抓包

2026/9/8 4:55:53

前言 本文介绍如何通过 tcpdump 在 Android 手机上抓取网络数据包,并在电脑端使用 Wireshark 进行分析。适用于需要排查 App 网络请求、分析接口调用或调试网络问题的开发与测试场景。1. 手机要有 root 权限2. 下载 tcpdump3. adb push C:\Users\zhangkuixun\Downlo…

大模型推理镜像极简瘦身:从 25GB 巨无霸到 3GB 精简镜像实战

大模型推理镜像极简瘦身:从 25GB 巨无霸到 3GB 精简镜像实战

2026/9/8 22:37:26

大模型推理镜像极简瘦身:从 25GB 巨无霸到 3GB 精简镜像实战 在云原生基础设施中,容器镜像体积直接决定了服务的部署速度与弹性扩容敏捷度。对于传统的 Go / Java 微服务,镜像体积通常被严格控制在 50MB 到 200MB 以内,拉取镜像只…

扩散模型图像恢复实战:从DDPM原理到PyQt5可视化系统

扩散模型图像恢复实战:从DDPM原理到PyQt5可视化系统

2026/9/9 0:03:36

简介:面向毕业设计场景的PyQt5扩散模型图像恢复项目,提供完整Python源码与项目说明,适合图像处理、深度学习方向的高年级本科生与研究生参考。项目在模块设计上覆盖图像处理、扩散模型、参数配置、用户界面与结果评估五部分,具体涉…

开关电源环路裕量测试实战:相位裕量与增益裕量详解

开关电源环路裕量测试实战:相位裕量与增益裕量详解

2026/9/9 0:03:36

1. 项目概述:为什么环路裕量测试是电子工程师绕不开的“体检项目”“从零开始的电子工程师生活(6)——环路裕量测试”,这个标题一出来,老电源工程师可能已经下意识摸了摸示波器探头,新同事则大概率在想&…

定时插座芯片怎么选?专用定时IC与单片机MCU选型对比

定时插座芯片怎么选?专用定时IC与单片机MCU选型对比

2026/9/9 0:03:36

拆开市面上不同价位的定时插座,你会发现一个有意思的现象:有的里面躺着一颗黑色的软封装芯片,丝印都看不清;有的则是一块小小的蓝色或绿色PCB,上面赫然印着STM8或者STC的字样。同样叫"定时插座",…

远程协作的工作台整理

远程协作的工作台整理

2026/9/9 16:28:52

远程协作的工作台整理远程协作的核心不是再加一个工具,而是让交接信息足够完整。异步任务要写明目标、输入位置、完成标准和需要决策的人。 工作台的最小配置 将日程、待办、代码和沟通入口收拢到少数固定位置;通知按紧急程度分层。工作台不需要模仿办公…

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能

2026/9/8 3:19:39

持续集成 流水线自动化与 声明式交付 实践:原型怎样变成可用功能分类:[AI/大模型]细分主题:AI 增强型 CI/CD 流水线自动化与 GitOps 实践:Agent 工作流、工具调用与任务拆解:从原型到生产的验收清单很多团队在尝试用大…

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场

2026/9/8 4:00:23

容器编排 生产环境运维与排障实战:复盘记录怎样真正派上用场分类:[工程技术]细分主题:Kubernetes 生产环境运维与排障实战:可复制的项目复盘模板与决策记录大部分团队的事故复盘报告,最后都变成了躺在 Confluence 或钉…