跳到主要内容

创建 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.job1
    job1-worker-0.job1
    job1-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()