首页
/ Colossal-AI 分布式优化器全解析:Adafactor、CAME、GaLore 与 LAMB 如何无缝对接 Tensor Parallel 与 ZeRO

Colossal-AI 分布式优化器全解析:Adafactor、CAME、GaLore 与 LAMB 如何无缝对接 Tensor Parallel 与 ZeRO

2026-09-08 12:02:11作者:谭伦延

本文基于 Colossal-AI 官方中文文档 distributed_optimizers.md 并结合仓库源码,系统讲解 Colossal-AI 内置的四种内存高效分布式优化器:Adafactor、CAME、GaLore 与 LAMB。文章将说明为何这些需要"逐层统计信息"的优化器无法直接用于模型并行场景,展示它们如何通过 booster 与 HybridParallelPlugin 在 Tensor Parallel + ZeRO 下训练,并给出可直接运行的完整代码、关键参数表与插件兼容性矩阵,让你拿到手就能在 Colossal-AI 中落地训练。

为什么需要"分布式优化器"

在分布式训练中,常用手段是把模型切分到多个设备上:Tensor Parallel(TP)将单个算子/权重按张量维度切分到一组设备,ZeRO(Zero Redundancy Optimizer)则将优化器状态、梯度乃至参数按数据并行维度分片。Adam 与 SGD 这类经典优化器(update 只依赖当前梯度的逐元素运算)可以很自然地在分片视图上运行。

但 Adafactor、CAME、GaLore、LAMB 这一类现代优化器完全不同:它们为了在降低内存稳定大 batch 训练的同时保持收敛,普遍依赖逐参数/逐层的统计信息——例如 Adafactor 需要对每行的二阶矩做行均值、列均值分解,LAMB 需要逐层计算"trust ratio"。当模型层被切到不同设备上(TP),或者参数被 ZeRO 按行分片之后,这些统计量的归约、重建和低秩投影就变得棘手——这就是 Colossal-AI 提供 DistributedAdaFactorDistributedLAMBDistGaloreAwamWDistributedCAME 等分布式版本优化器的根本原因。

好消息是:即使你不主动实例化分布式优化器,Colossal-AI 的 plugin 也会在 booster.boost 时自动把普通优化器转换为对应的分布式版本,真正做到开箱即用。

四种优化器:原理速览

文档在"优化器"一节对四种算法做了高度概括,结合源码可以展开如下:

优化器 核心思想 内存收益来源 对应论文
Adafactor 首个引入非负矩阵分解(NMF)的 Adam 变体,把逐元素的二阶矩张量 [H, W] 近似为"行向量 × 列向量"的低秩形式 row * col 二阶矩存储从 O(H×W) 降到 O(H+W) Adafactor: Adaptive Learning Rates with Sublinear Memory Cost(arXiv:1804.04235)
CAME 在 Adafactor 的低秩分解之上引入**置信度矩阵(confidence matrix)**与三个 beta 系数的动量更新,缓解低秩近似带来的"过度置信"问题 保持低秩二阶矩,仅增加一组轻量置信度状态 CAME: Confidence-guided Adaptive Memory Efficient Optimization(arXiv:2307.02047)
GaLore 将梯度通过周期性 SVD 投影到低秩子空间进行更新,配合 8-bit 块状量化进一步压内存 梯度与优化器状态都"瘦身"为低秩形态 GaLore: Memory-Efficient LLM Training by Gradient Low-Rank Projection(arXiv:2403.03507)
LAMB 通过以 Lipschitz 常数倒数界定的逐层自适应更新(layer-wise adaptive update),在超大批量(如 BERT 76 分钟训练)下不失精度 本身并非省内存,而是允许超大 batch 以摊薄同步开销、缩短训练时间 Large Batch Optimization for Deep Learning: Training BERT in 76 minutes(arXiv:1904.00962)

从代码结构看,Colossal-AI 在 colossalai/nn/optimizer/ 目录中为每一种算法都准备了两套实现:

  • 单机可用的普通实现:adafactor.pycame.pygalore.pylamb.py
  • 面向 TP/ZeRO 的分布式实现:distributed_adafactor.pydistributed_came.pydistributed_galore.pydistributed_lamb.py

自动转换机制:plugin 内部做了什么

原文档明确写道:"即使您不使用 distributed optimizer,plugin 也会自动将 optimizer 转换为分布式版本以方便使用。"源码印证了这一设计。在 colossalai/nn/optimizer/init.py 中定义了普通优化器到分布式版本的注册表和转换函数:

optim2DistOptim = {
    GaLoreAdamW8bit: DistGaloreAwamW,
    Lamb: DistributedLamb,
    CAME: DistributedCAME,
    Adafactor: DistributedAdaFactor,
}

def cast_to_distributed(optim):
    if optim.__class__ in optim2DistOptim:
        _logger = get_dist_logger()
        _logger.info(f"Converting optimizer {optim.__class__.__name__} to its distributed version.", ranks=[0])
        if isinstance(optim, GaLoreAdamW8bit):
            return optim2DistOptimGaLoreAdamW8bit
        return optim2DistOptimoptim.__class__
    return optim

转换发生后,第 0 号 rank 会打印一条 Converting optimizer X to its distributed version. 的日志,可以作为排查"我的优化器到底是不是分布式版本"的依据。实际调用 cast_to_distributed 的 plugin 包括:

这些分布式优化器共同继承自 colossalai/interface/optimizer.py 中定义的分布式优化器接口,并通过 setup_distributed(tp_group, dp_group, shard_to_working_param, padding_map, use_zero/is_zero) 方法注入 Tensor Parallel 与数据并行(ZeRO)的通信组。也就是说,plugin 拿到分布式优化器后,会负责把正确的进程组和分片信息"喂"给它,用户无需自己管理 all_reduce/gather/split

五步上手:用 Adafactor 跑通 TP + ZeRO 2

下面完整展开原文档中的使用流程。示例以 HybridParallelPlugin(tp_size=2, zero_stage=2) 为例,将 tp_size=2 的 Tensor Parallel 与 ZeRO stage 2 叠加使用。

Step 1:导入依赖

from transformers import LlamaModel, LlamaConfig
from colossalai.nn.optimizer.distributed_adafactor import DistributedAdaFactor
from colossalai.booster import Booster
from colossalai.booster.plugin import HybridParallelPlugin
import colossalai
import torch

Step 2:初始化分布式环境

colossalai.launch_from_torch()

launch_from_torch() 适用于由 torchrun/colossal run 拉起的环境。原文档在注释中给出了推荐的启动方式 colossalai run --nproc_per_node 4(文档末尾还附带了自动化 doc-test 命令 colossalai run --nproc_per_node 4 distributed_optimizers.py,保证示例可被 CI 直接复跑)。更多初始化方式(colossalai.launchlaunch_from_slurm 等)可参考 Launch Colossal-AI

Step 3:初始化模型与分布式优化器

configuration = LlamaConfig()
model = LlamaModel(configuration).cuda()
criterion = lambda x: x.mean()
dist_optim = DistributedAdaFactor(model.parameters())

这里直接实例化了分布式版本 DistributedAdaFactor。注意它的默认参数与 Hugging Face 的 Adafactor 对齐:lr=Noneeps=(1e-30, 1e-3)clip_threshold=1.0decay_rate=-0.8beta1=Noneweight_decay=0.0scale_parameter=Truerelative_step=Truewarmup_init=False

  • relative_step=True(默认)时,学习率由更新步数自动推算:min(1e-2, 1/sqrt(step))(若开启 warmup_init 则下限为 1e-6*step),此时不允许同时传入手动 lr,源码中会抛出 ValueError: Cannot combine manual lr and relative_step=True options
  • warmup_init=True 强制要求 relative_step=True
  • 默认 beta1=None 表示不启用一阶矩(进一步省内存);给 beta1 赋值后才会维护 exp_avg 状态;
  • 对于维度 ≥ 2 的参数张量,算法采用 factored 方式存储行/列统计量 exp_avg_sq_row(形状 [H])与 exp_avg_sq_col(形状 [W]);一维参数(如 bias)则退化为完整的 exp_avg_sq。这一点由 distributed_adafactor.py_get_optionsfactored = len(param_shape) >= 2 决定。

Step 4:初始化 plugin 与 booster

plugin = HybridParallelPlugin(tp_size=2, zero_stage=2, pp_size=1, enable_all_optimization=True)
booster = Booster(plugin=plugin)
# You should also pass in your own dataset.
model, dist_optim, criterion, dataloader, _ = booster.boost(model, dist_optim, criterion)

booster.boost 是整套机制的中枢:它根据 plugin 配置对模型做切分(Tensor Parallel 经 Shardformer 完成)、把参数交给 ZeRO 2 分片管理,并在内部调用 setup_distributed 为分布式优化器注入正确的 tp_group/dp_groupshard_to_working_param 映射。shard_to_working_param 是 ZeRO 下参数"分片视图"到前反向实际使用的工作参数之间的桥接,分布式优化器正是借助它判断当前参数是否为分布式张量(is_distributed_tensor)、读取其 ShardingSpec(get_sharding_spec),从而决定如何归约统计量。

distributed_adafactor.pysetup_distributed 实现可以看到它对不同并行切分的精细化处理:

  • 列并行权重(Col Parallel),即 ShardingSpec 首维为 R:沿行方向的二阶矩 exp_avg_sq_row 需要在 tp_group 上做 all_reduce 再取平均;
  • 行并行权重(Row Parallel),即 ShardingSpec 末维为 R:列方向的统计量需要在 TP 组内归约,必要时还需 gather 完整行后求均值;
  • 配合 ZeRO 时,RMS(均方根)计算还会通过 dist.all_reducedp_group 上把分片统计量归约完整,计算方式见 _rms 静态方法。

这正是"分布式优化器"区别于朴素实现的核心价值:所有跨设备的归约都按分片布局自动完成,且做到正确性与单机等价。

Step 5:训练循环

steps = 10
for step in range(steps):
    input_ids = torch.ones(1, 100, device="cuda", dtype=torch.int)
    attention_mask = input_ids.clone()
    outputs = model(input_ids.cuda(), attention_mask.cuda())
    loss = criterion(outputs.last_hidden_state)
    booster.backward(loss, dist_optim)
    dist_optim.step()
    dist_optim.zero_grad()

注意:step() 内部(对应 Adafactor 论文算法)在更新时先计算 beta2t = 1 - step^decay_rate(decay_rate 默认 -0.8),随后做 RMS 裁剪——update.div_((rms / clip_threshold).clamp_(min=1.0)),再用 relative step 学习率缩放;scale_parameter=True 时学习率还会乘上 max(eps[1], RMS) 作为参数尺度校正。这套逻辑在单机 adafactor.py 与分布式版中保持一致,区别仅在于分布式版多了通信归约。

GaLore 的特殊配置:低秩投影 + 8-bit 量化

GaLore 与其余三者最大的不同在于:它需要为每个参数组显式指定投影秩(rank),并可叠加 bitsandbytes 的 8-bit 量化与分页(paged)优化器。原文档提示:量化的详细行为可参考 bitsandbytes 的实现。

参数分组:get_galore_param_groups

GaLore 只对二维矩阵参数(如 nn.Linear 的权重、词嵌入表)做低秩 SVD 投影才有意义。因此仓库提供了工具函数 get_galore_param_groups,它遍历模型自动分成三组:

def get_galore_param_groups(model, weight_decay, rank=256, update_proj_gap=200, scale=0.25, proj_type="std"):
    # ...
    no_decay = ["bias", "LayerNorm.weight"]

分组规则为:

  1. galore_params:维度等于 2 的参数(param.dim() == 2),应用 rank 低秩投影;
  2. non_galore:其余非二维参数,走普通 AdamW 更新;
  3. no_decay_params:名字包含 biasLayerNorm.weight 的参数,weight_decay 置 0。

分布式 GaLore 实例化

文档给出的分布式 GaLore(8-bit AdamW 底座)完整配置如下(此处按当前仓库接口将分组函数的衰减参数名统一为 weight_decay):

from colossalai.nn.optimizer.galore import get_galore_param_groups
from colossalai.nn.optimizer import DistGaloreAwamW

optim = DistGaloreAwamW(
    get_galore_param_groups(model, weight_decay=1e-2, rank=8),
    lr=lr,
    betas=(beta1, beta2),
    eps=eps,
    nbits=8,
    percentile_clipping=100,
    block_wise=True,
    min_8bit_size=4096,
)

对照 distributed_galore.py 的构造签名,关键参数含义如下:

参数 默认值 说明
rank(分组级) 256 低秩投影的秩,越大保留信息越多、省内存越少
update_proj_gap(分组级) 200 每隔多少步重做一次 SVD 刷新正交投影矩阵
scale(分组级) 0.25 低秩回投(project_back)时的缩放因子
proj_type(分组级) "std" 投影方向,可参考 galore_torchstd/reverse_std 语义
nbits 8 优化器状态量化位数,仅支持 32 与 8
min_8bit_size 4096 参数元素数低于该值时不做 8-bit 优化
percentile_clipping 100 按梯度范数百分位自适应裁剪阈值,提升稳定性
block_wise True 分块独立量化,缓解离群值对量化的影响
is_paged False 是否为 paged 优化器(通过 CPU-GPU 换入换出应对显存尖峰)

DistGaloreAwamW 同时继承了 Colossal-AI 的 DistributedOptim 与 bitsandbytes 的 Optimizer2State,因此它既获得了 TP/ZeRO 分布式注入能力,又复用 bitsandbytes 的 8-bit 优化内核。构造完成后,若所有参数组都没有 rank 字段,会打印警告 "Will not apply GaLore as rank isn't in any param group...",提醒你改用 get_galore_param_groups

其底层原理由 galore.pyGaLoreProjector 承载:维护一个周期更新的正交矩阵,每次迭代先 project() 把满秩梯度投影成低秩梯度参与更新,再 project_back() 乘回 scale 恢复满秩形状——整套机制对优化器状态实现了约 rank/min(m,n) 比例的压缩。在 ZeRO 2 场景下,低秩投影是在分片梯度视图上完成的,因此 setup_distributed 要求 ZeRO 提供每个参数的 padding 信息(padding_map)以保证 SVD 的正确性。

其余优化器的关键参数速查

DistributedLAMB

distributed_lamb.py。默认 lr=1e-3betas=(0.9, 0.999)eps=1e-6weight_decay=0bias_correction=True,构造时会校验学习率、eps 与 betas 的取值合法性。LAMB 的核心是按层计算更新量与参数范数的 trust ratio,因此一旦参数被 TP 切分或 ZeRO 分片,就必须对每层范数做跨设备归约——这正是 DistributedLamb 提供的增量能力。适合超大 batch 训练,通常配合较大学习率使用。

DistributedCAME

distributed_came.py。默认 lr=Noneeps=(1e-30, 1e-16)clip_threshold=1.0betas=(0.9, 0.999, 0.9999)weight_decay=0.0。与 Adafactor 不同,CAME 使用三个 beta 分别跟踪 update、平方梯度与"不稳定性/置信度"项,这就是论文中 confidence-guided 的来源。类上标注 supports_memory_efficient_fp16 = True,表示可在混合精度下开启省内存的 FP16 路径。它的 setup_distributed 还针对分布式场景做了防御性处理:当张量在行/列并行下被切到 H=1 这类极端形状时,强制将其判为 factored(factored=True),避免因维度退化为 1 而误走非分解路径导致内存爆炸。

与 Plugin 的兼容性矩阵

原文档给出了官方兼容性矩阵,四种分布式优化器在五个插件下的支持情况如下(✔️ 支持,❌ 不支持):

Optimizer / Plugin Hybrid Parallel Plugin Low Level Zero Plugin Torch DDP Plugin Gemini Plugin MoE Hybrid Plugin
LAMB ✔️ ✔️ ✔️
GaLore ✔️ ✔️ ✔️
Adafactor ✔️ ✔️ ✔️
CAME ✔️ ✔️ ✔️

读取该表需要注意两点:

  1. 前三列(Hybrid Parallel、Low Level Zero、Torch DDP)是完整的实践组合——Hybrid Parallel 插件通常意味着 TP + ZeRO/DP 混布,这正是文档 Step 4 场景;
  2. Gemini(Gemini Plugin 采用异构内存 + chunk 式 ZeRO 管理)与 MoE Hybrid 插件目前不参与这四种分布式优化器的兼容矩阵,选择组合时需要避开;
  3. 兼容性是会演进的:例如 MoE Hybrid 插件的源码中已出现 cast_to_distributed 调用点(见上文),后续版本的支持范围请以当前仓库代码与官方文档为准。

测试与验证

仓库在 tests/test_optimizer/ 目录下为四种分布式优化器提供了专门的测试用例:

  • test_dist_adafactor.py
  • test_dist_came.py
  • test_dist_galore.py
  • test_dist_lamb.py

这些测试在分布式环境下实际创建进程组并跑优化步骤,用于校验分片梯度下的统计量归约与参数更新结果,可作为你参考或复跑验证的起点。若要快速验证本文示例,可先起一个小规模 TP 配置(如 tp_size=1, zero_stage=1)单机验证逻辑,再逐步放大到多卡。

小结

总结起来,在 Colossal-AI 中使用 Adafactor/CAME/GaLore/LAMB 这类统计量敏感的优化器,遵循一条最省事的路径即可:

  1. 普通写法:直接用 AdaFactor/Lamb/CAME/GaLoreAdamW8bit,或 GaLore 用 get_galore_param_groups + DistGaloreAwamW 手工指定 rank;
  2. 交给 booster:HybridParallelPlugin(TP + ZeRO)等插件会在 booster.boost 阶段通过 optim2DistOptim 自动完成普通优化器 → 分布式优化器的转换并注入进程组;
  3. 唯一需要你操心的只有两件事:GaLore 的 rank 与量化参数、以及插件兼容性选择(避开 Gemini/MoE Hybrid)。

通过这种"注册表自动转换 + plugin 注入通信组"的架构,Colossal-AI 把最复杂的跨设备统计量归约隐藏在了框架内部,让研究者既能享受低秩/分解类优化器的内存红利,又不必手写任何分布式归约逻辑。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
898
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
921
1.84 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.8 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
531
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
519
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
391