创建 Colossal-AI 分布式训练
更新时间:2025-08-18 15:00:32
Colossal-AI 是一个旨在使大型 AI 模型的训练、微调和推理变得更便宜、更快速、更容易获得的开源深度学习系统。它的核心优势在于其强大的并行化功能,它集成并统一了多种并行策略,包括数据并行、流水线并行、多维度张量并行以及零冗余优化器 (ZeRO) 等,让开发者可以用编写单卡程序的方式轻松实现复杂的分布式训练。
前提条件
- 选择的镜像中需要已安装 ColossalAI、openssh-server。
操作步骤
请参考创建分布式训练模块填写基本配置,计算框架选择 Colossal-AI。
任务配置中固定有2个 Role(节点组):
- master:只有1个节点,执行发起 colossalai 训练的节点,通过训练任务节点间免密 SSH 通信在其他 worker 节点同步启动训练。
- worker:可设置多个节点,每个 worker 节点仅启动 sshd 服务,等待 master 节点发起训练。
在 master 节点输入启动命令示例如下:
colossalai run \
--nproc_per_node 8 \
--hostfile /etc/colossalai/hostfile \
<client_entry.py> \
<client args>
其中
-
--hostfile文件中记录了当前 master 和 worker 所有节点地址,请使用固定文件/etc/colossalai/hostfile,改文件由平台生成内置。文件内容示例如下:job1-master-0.job1job1-worker-0.job1job1-worker-1.job1 -
<client_entry.py>是您的启动脚本。 -
<client args>是你启动脚本接收的参数。
更多 DeepSpeed 使用方法请参考 DeepSpeed 官方文档:
代码示例
import argparse
import os
from pathlib import Path
import torch.nn as nn
import torchvision
import torchvision.transforms as transforms
from torch.optim import Optimizer
from torch.optim.lr_scheduler import MultiStepLR
from torch.utils.data import DataLoader
from tqdm import tqdm
import colossalai
from colossalai.booster import Booster
from colossalai.booster.plugin import GeminiPlugin, LowLevelZeroPlugin, TorchDDPPlugin
from colossalai.booster.plugin.dp_plugin_base import DPPluginBase
from colossalai.cluster import DistCoordinator
from colossalai.nn.optimizer import HybridAdam
from colossalai.utils import get_current_device
NUM_EPOCHS = 80
LEARNING_RATE = 1e-3
def build_dataloader(batch_size: int, coordinator: DistCoordinator, plugin: DPPluginBase):
transform_train = transforms.Compose(
[transforms.Pad(4), transforms.RandomHorizontalFlip(), transforms.RandomCrop(32), transforms.ToTensor()])
transform_test = transforms.ToTensor()
data_path = os.environ.get("DATA", "./data")
with coordinator.priority_execution():
train_dataset = torchvision.datasets.CIFAR10(
root=data_path, train=True, transform=transform_train, download=True)
test_dataset = torchvision.datasets.CIFAR10(
root=data_path, train=False, transform=transform_test, download=True)
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)
return train_dataloader, test_dataloader
def train_epoch(
epoch: int,
model: nn.Module,
optimizer: Optimizer,
criterion: nn.Module,
train_dataloader: DataLoader,
booster: Booster,
coordinator: DistCoordinator,
):
model.train()
with tqdm(train_dataloader, desc=f"Epoch [{epoch + 1}/{NUM_EPOCHS}]", disable=not coordinator.is_master()) as pbar:
for images, labels in pbar:
images = images.cuda()
labels = labels.cuda()
# Forward pass
outputs = model(images)
loss = criterion(outputs, labels)
# Backward and optimize
booster.backward(loss, optimizer)
optimizer.step()
optimizer.zero_grad()
# Print log info
pbar.set_postfix({"loss": loss.item()})
def main():
parser = argparse.ArgumentParser()
parser.add_argument(
"-p",
"--plugin",
type=str,
default="torch_ddp",
choices=["torch_ddp", "torch_ddp_fp16", "low_level_zero", "gemini"],
help="plugin to use",
)
parser.add_argument("-r", "--resume", type=int, default=-1, help="resume from the epoch's checkpoint")
parser.add_argument("-c", "--checkpoint", type=str, default="./checkpoint", help="checkpoint directory")
parser.add_argument("-i", "--interval", type=int, default=5, help="interval of saving checkpoint")
parser.add_argument("--target_acc", type=float, default=None, help="target accuracy. Raise exception if not reached")
args = parser.parse_args()
if args.interval > 0:
Path(args.checkpoint).mkdir(parents=True, exist_ok=True)
colossalai.launch_from_torch(config={})
coordinator = DistCoordinator()
global LEARNING_RATE
LEARNING_RATE *= coordinator.world_size
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(initial_scale=2**5)
elif args.plugin == "low_level_zero":
plugin = LowLevelZeroPlugin(initial_scale=2**5)
booster = Booster(plugin=plugin, **booster_kwargs)
train_dataloader, test_dataloader = build_dataloader(100, coordinator, plugin)
model = torchvision.models.resnet18(num_classes=10)
criterion = nn.CrossEntropyLoss()
optimizer = HybridAdam(model.parameters(), lr=LEARNING_RATE)
lr_scheduler = MultiStepLR(optimizer, milestones=[20, 40, 60, 80], gamma=1 / 3)
model, optimizer, criterion, _, lr_scheduler = booster.boost(
model, optimizer, criterion=criterion, lr_scheduler=lr_scheduler
)
if args.resume >= 0:
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")
start_epoch = args.resume if args.resume >= 0 else 0
for epoch in range(start_epoch, NUM_EPOCHS):
train_epoch(epoch, model, optimizer, criterion, train_dataloader, booster,coordinator)
lr_scheduler.step()
if args.interval > 0 and (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")
if __name__ == "__main__":
main()