Colossal-AI 分布式优化器全解析:Adafactor、CAME、GaLore 与 LAMB 如何无缝对接 Tensor Parallel 与 ZeRO
本文基于 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 提供 DistributedAdaFactor、DistributedLAMB、DistGaloreAwamW、DistributedCAME 等分布式版本优化器的根本原因。
好消息是:即使你不主动实例化分布式优化器,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.py、came.py、galore.py、lamb.py; - 面向 TP/ZeRO 的分布式实现:
distributed_adafactor.py、distributed_came.py、distributed_galore.py、distributed_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 包括:
- hybrid_parallel_plugin.py:在配置 optimizer 阶段自动转换;
- low_level_zero_plugin.py:Low Level ZeRO 插件同样转换;
- moe_hybrid_parallel_plugin.py:MoE 混合并行插件也接入了该机制。
这些分布式优化器共同继承自 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.launch、launch_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=None、eps=(1e-30, 1e-3)、clip_threshold=1.0、decay_rate=-0.8、beta1=None、weight_decay=0.0、scale_parameter=True、relative_step=True、warmup_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_options的factored = 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_group 与 shard_to_working_param 映射。shard_to_working_param 是 ZeRO 下参数"分片视图"到前反向实际使用的工作参数之间的桥接,分布式优化器正是借助它判断当前参数是否为分布式张量(is_distributed_tensor)、读取其 ShardingSpec(get_sharding_spec),从而决定如何归约统计量。
从 distributed_adafactor.py 的 setup_distributed 实现可以看到它对不同并行切分的精细化处理:
- 列并行权重(Col Parallel),即 ShardingSpec 首维为
R:沿行方向的二阶矩exp_avg_sq_row需要在tp_group上做all_reduce再取平均; - 行并行权重(Row Parallel),即 ShardingSpec 末维为
R:列方向的统计量需要在 TP 组内归约,必要时还需gather完整行后求均值; - 配合 ZeRO 时,RMS(均方根)计算还会通过
dist.all_reduce在dp_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"]
分组规则为:
- galore_params:维度等于 2 的参数(
param.dim() == 2),应用rank低秩投影; - non_galore:其余非二维参数,走普通 AdamW 更新;
- no_decay_params:名字包含
bias或LayerNorm.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_torch 的 std/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.py 的 GaLoreProjector 承载:维护一个周期更新的正交矩阵,每次迭代先 project() 把满秩梯度投影成低秩梯度参与更新,再 project_back() 乘回 scale 恢复满秩形状——整套机制对优化器状态实现了约 rank/min(m,n) 比例的压缩。在 ZeRO 2 场景下,低秩投影是在分片梯度视图上完成的,因此 setup_distributed 要求 ZeRO 提供每个参数的 padding 信息(padding_map)以保证 SVD 的正确性。
其余优化器的关键参数速查
DistributedLAMB
见 distributed_lamb.py。默认 lr=1e-3、betas=(0.9, 0.999)、eps=1e-6、weight_decay=0、bias_correction=True,构造时会校验学习率、eps 与 betas 的取值合法性。LAMB 的核心是按层计算更新量与参数范数的 trust ratio,因此一旦参数被 TP 切分或 ZeRO 分片,就必须对每层范数做跨设备归约——这正是 DistributedLamb 提供的增量能力。适合超大 batch 训练,通常配合较大学习率使用。
DistributedCAME
见 distributed_came.py。默认 lr=None、eps=(1e-30, 1e-16)、clip_threshold=1.0、betas=(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 | ✔️ | ✔️ | ✔️ | ❌ | ❌ |
读取该表需要注意两点:
- 前三列(Hybrid Parallel、Low Level Zero、Torch DDP)是完整的实践组合——Hybrid Parallel 插件通常意味着 TP + ZeRO/DP 混布,这正是文档 Step 4 场景;
- Gemini(Gemini Plugin 采用异构内存 + chunk 式 ZeRO 管理)与 MoE Hybrid 插件目前不参与这四种分布式优化器的兼容矩阵,选择组合时需要避开;
- 兼容性是会演进的:例如 MoE Hybrid 插件的源码中已出现
cast_to_distributed调用点(见上文),后续版本的支持范围请以当前仓库代码与官方文档为准。
测试与验证
仓库在 tests/test_optimizer/ 目录下为四种分布式优化器提供了专门的测试用例:
test_dist_adafactor.pytest_dist_came.pytest_dist_galore.pytest_dist_lamb.py
这些测试在分布式环境下实际创建进程组并跑优化步骤,用于校验分片梯度下的统计量归约与参数更新结果,可作为你参考或复跑验证的起点。若要快速验证本文示例,可先起一个小规模 TP 配置(如 tp_size=1, zero_stage=1)单机验证逻辑,再逐步放大到多卡。
小结
总结起来,在 Colossal-AI 中使用 Adafactor/CAME/GaLore/LAMB 这类统计量敏感的优化器,遵循一条最省事的路径即可:
- 普通写法:直接用
AdaFactor/Lamb/CAME/GaLoreAdamW8bit,或 GaLore 用get_galore_param_groups+DistGaloreAwamW手工指定 rank; - 交给 booster:
HybridParallelPlugin(TP + ZeRO)等插件会在booster.boost阶段通过optim2DistOptim自动完成普通优化器 → 分布式优化器的转换并注入进程组; - 唯一需要你操心的只有两件事:GaLore 的 rank 与量化参数、以及插件兼容性选择(避开 Gemini/MoE Hybrid)。
通过这种"注册表自动转换 + plugin 注入通信组"的架构,Colossal-AI 把最复杂的跨设备统计量归约隐藏在了框架内部,让研究者既能享受低秩/分解类优化器的内存红利,又不必手写任何分布式归约逻辑。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0629
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
HivisionIDPhotos⚡️HivisionIDPhotos: a lightweight and efficient AI ID photos tools. 一个轻量级的AI证件照制作算法。Python07
DragonOSDragonOS is an operating system developed from scratch using Rust, with Linux compatibility. It is designed for **Serverless** scenarios. 使用Rust从0自研内核,具有Linux兼容性的操作系统,面向云计算Serverless场景而设计。Rust00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00