ColossalAI 中的 PaLM 教学实现:不到 200 行核心代码与 Booster 分布式训练实战
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
导读
本篇文章基于 ColossalAI 仓库内的examples/language/palm示例(关联文档为 examples/language/palm/README.md),系统梳理 Google PaLM(Pathways Language Model)自回归 Transformer 的教学实现,深入讲解其不足 200 行的核心模型代码,并结合仓库内的 train.py、test_ci.sh 等文件,说明如何借助 ColossalAI 的 Booster API 在 Enwik8 数据集上完成单机多卡训练与吞吐测试。读完本文,你将掌握 PaLM 架构的关键设计(并行残差、RoPE、SwiGLU、Multi-Query Attention、无偏置 LayerNorm)以及 ColossalAI 插件化训练(TorchDDP / Gemini / Low-Level Zero)的完整实操方法。
一、示例背景:教育用途的 PaLM 精简复现
examples/language/palm示例复现了 Google 于 2022 年发表的《PaLM: Scaling Language Modeling with Pathways》中所描述的 Transformer 架构。整个核心模型实现在 palm_pytorch/palm_pytorch.py 中,正文加上注意力、前馈等子模块,整体依然保持精简——正如原文档所述,这是一个"用于教学目的、说明大规模语言模型核心其实并不复杂"的实现。论文引用信息也一并收录在原文档的 Citations 小节中(Chowdhery, Aakanksha et al., 2022)。
需要说明的是:本示例并非为了复现 540B 参数的极致扩展性,而是为了揭示架构本质。原 README 对"是否 SOTA"、"能否扩展到数千亿参数"等说法持保留与教育导向的态度,这一点在成文时同样予以继承。
二、模型结构逐模块拆解:不到 200 行的 PaLM
核心模型定义于 palm_pytorch/palm_pytorch.py,顶层通过函数PaLM(*, dim, num_tokens, depth, dim_head=64, heads=8, ff_mult=4)工厂式构建(见 palm_pytorch.py 中 PaLM 构建段)。整体是一个nn.Sequential:先是 token Embedding,随后堆叠depth个ParallelResidual(Attention, FeedForward)层,最后是 LayerNorm 与输出投影 Linear。下面按实现顺序拆解各模块的设计意图。
2.1 无偏置 LayerNorm:刻意为之的复现细节
class LayerNorm(nn.Module): def __init__(self, dim, eps=1e-5): super().__init__() self.eps = eps self.gamma = nn.Parameter(torch.ones(dim)) self.register_buffer("beta", torch.zeros(dim)) def forward(self, x): return F.layer_norm(x, x.shape[-1:], self.gamma, self.beta)源码注释说明:PaLM 使用不含偏置(bias)的 LayerNorm,而 PyTorch 原生的nn.LayerNorm默认含 bias,因此这里手动实现一个gamma参数 +beta零缓冲区的版本,以满足复现要求。
2.2 ParallelResidual:并行分支 + 残差连接
class ParallelResidual(nn.Module): def __init__(self, *fns): super().__init__() self.fns = nn.ModuleList(fns) def forward(self, x): return x + sum([fn(x) for fn in self.fns])ParallelResidual把注意力与前馈两个分支并行地作用在同一份输入x上,再把输出直接相加后叠加上残差x。源码注释指出这一模式由 Wang 等人与 EleutherAI(GPT-J 时期)发现,也是 PaLM 与常见串行"Pre-Norm Transformer"的一大区别——两个子层共享同一个输入与同一条残差路径,能减少序列化依赖、提升硬件利用率。
2.3 RotaryEmbedding:旋转位置编码
class RotaryEmbedding(nn.Module): def __init__(self, dim): super().__init__() inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq) def forward(self, max_seq_len, *, device): seq = torch.arange(max_seq_len, device=device) i, j = len(seq.type_as(self.inv_freq)), len(self.inv_freq) freqs = matmul(seq.type_as(self.inv_freq).reshape(i, 1), self.inv_freq.reshape(1, j)) return torch.cat((freqs, freqs), dim=-1)位置编码采用 RoPE 方案:按频率1 / 10000^(2i/dim)构造旋转角,与序列位置矩阵相乘后复制拼接。查询q与键k都会在注意力之前通过apply_rotary_pos_emb(pos, t)施加旋转位置信息,从而在不改动绝对坐标的前提下把相对位置关系编码进点积中。Attention内部还通过register_buffer("pos_emb", None, persistent=False)对掩码与位置编码做按需缓存(get_mask/get_rotary_embedding仅在序列变长时才重算),减少重复计算。
2.4 SwiGLU 前馈网络
class SwiGLU(nn.Module): def forward(self, x): x, gate = x.chunk(2, dim=-1) return F.silu(gate) * x def FeedForward(dim, mult=4): inner_dim = int(dim * mult) return nn.Sequential( LayerNorm(dim), nn.Linear(dim, inner_dim * 2, bias=False), SwiGLU(), nn.Linear(inner_dim, dim, bias=False), )前馈部分把输入升维到inner_dim * 2,拆成两半,一半经F.silu(即 SiLU/Swish)激活后与另一半逐元素相乘,形成 SwiGLU 门控。源码注释点明:这是 Noam Shazeer 论文中的经典做法,只是本实现选择 SwiGLU 而非更流行的 GEGLU。注意每个前馈模块开头同样先过无偏置 LayerNorm。
2.5 Attention:Multi-Query 键值 + 因果掩码
Attention的核心设计有几点:
self.to_q将输入投影为heads * dim_head维的 Q;self.to_kv只投影出一份dim_head * 2维的 K/V——这正是Multi-Query / Multi-Head 但单键单值的注意力(源码注释指出这是 Shazeer 的另一篇论文思想:超过一定规模后无性能损失,且解码更高效);- Q 按头重排为
(b h n d),K/V 保持共享的单份; - 相似度计算、因果掩码填充(上三角掩码置为极小值)、减去行最大值后再 softmax(数值稳定技巧)均在源码中显式实现;
- V 聚合后重排回多头拼接,经
to_out投影输出。
2.6 顶层 PaLM:权重绑定与初始化
def PaLM(*, dim, num_tokens, depth, dim_head=64, heads=8, ff_mult=4): net = nn.Sequential( nn.Embedding(num_tokens, dim), *[ParallelResidual( Attention(dim=dim, dim_head=dim_head, heads=heads), FeedForward(dim=dim, mult=ff_mult), ) for _ in range(depth)], LayerNorm(dim), nn.Linear(dim, num_tokens, bias=False), ) net[-1].weight = net[0].weight # embedding 与输出投影共享权重 nn.init.normal_(net[0].weight, std=0.02) return net顶层把depth层并行残差块串起来;输出 Linear 的权重直接绑定回 Embedding 权重(weight tying),随后以标准差 0.02 的正态分布初始化词向量。其余默认超参数为dim_head=64, heads=8, ff_mult=4。
2.7 自回归封装:AutoregressiveWrapper
训练与生成由 palm_pytorch/autoregressive_wrapper.py 中的AutoregressiveWrapper承担:
forward:把输入拆成x[:, :-1]与标签x[:, 1:],让模型做标准的 next-token 预测,返回交叉熵损失;generate:自回归循环式生成。在torch.no_grad()与eval_decorator保护下逐 token 前向,取最后一个位置 logits 做top-k 过滤(filter_thres=0.9)并除以temperature后做 softmax,再torch.multinomial采样;支持eos_token提前终止与pad_value填充。
三、基础用法:从模型实例化到 540B 配置
原 README 给出的最简用法如下,读者可直接在安装依赖后运行验证:
import torch from palm_pytorch import PaLM palm = PaLM( num_tokens = 20000, dim = 512, depth = 12, heads = 8, dim_head = 64, ) tokens = torch.randint(0, 20000, (1, 2048)) logits = palm(tokens) # (1, 2048, 20000)输入为(batch=1, seq_len=2048)的 token 序列,输出形状为(1, 2048, 20000),即每个位置一个词表大小的 logits 分布。原文档同时也给出了论文中PaLM 540B对应的超参形态(仅作量级参考,本示例并不实际承载该规模):
palm = PaLM( num_tokens = 256000, dim = 18432, depth = 118, heads = 48, dim_head = 256, )四、使用 Booster 新 API 训练:Enwik8 实战
4.1 训练入口与命令行参数
原 README 明确指出:该示例在早期实现基础上接入了 ColossalAI 的Booster 新 API,以获得更灵活高效的训练方式与更好的易用性,训练入口为 train.py,并提供 test_ci.sh 脚本用于在多个插件下走通全流程。
通过 train.py 中 parse_args,可确认当前支持的命令行参数如下:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
--distplan | str | colossalai | 分布式方案,colossalai走 Booster 新 API,pytorch走原生 PyTorch 训练分支 |
--plugin | str | torch_ddp | Booster 插件,可选torch_ddp、torch_ddp_fp16、gemini、low_level_zero |
--offload_optim_frac | float | 1.0 | 优化器状态卸载比例,仅对 gemini 插件生效 |
--batch_size | int | 8 | 每个 DP 组的 batch size |
--dummy_data | bool | False | 是否使用随机整数构造的假数据 |
参数校验在模块加载阶段即完成:若--distplan不是colossalai/pytorch会直接抛出TypeError。
4.2 数据加载:Enwik8 与随机假数据
数据准备逻辑见 train.py 中 generate_dataset:
- 真实数据:从
./data/enwik8.gz读取(data/目录内的 README 说明 enwik8 数据源自 Hutter prize 页面),读取前 95e6 字节并按 90e6 / 5e6 切分为训练集与验证集,每字节按uint8词表处理; - 假数据:
--dummy_data=True时直接torch.randint(0, 100, ...)生成 9000 万 / 500 万长度的训练/验证序列,便于不下载数据即可进行 CI 与冒烟测试。
TextSamplerDataset(train.py)每次随机采样一段长度为seq_len + 1的连续字节作为一条样本(多出的 1 个 token 供自回归标签使用),并用cycle(DataLoader(...))无限循环供给训练循环。
4.3 ColossalAI 分支:模型、优化器与 Booster 装配
在--distplan=colossalai时(train.py):
model = PaLM(num_tokens=50304, dim=4096, depth=64) model = AutoregressiveWrapper(model, max_seq_len=SEQ_LEN) optimizer = HybridAdam(model.parameters(), lr=LEARNING_RATE, initial_scale=2**5) model, optimizer, _, _, _ = booster.boost(model, optimizer)该分支默认构造一个dim=4096、depth=64的较大模型以验证分布式插件的承载能力。几个值得注意的底层配合:
- 插件选择:
torch_ddp/torch_ddp_fp16使用TorchDDPPlugin()(fp16 时额外注入mixed_precision="fp16");gemini使用GeminiPlugin(offload_optim_frac=args.offload_optim_frac, initial_scale=2**5);low_level_zero使用LowLevelZeroPlugin(initial_scale=2**5); - LazyInitContext 与 Gemini 的配合:当插件为 gemini 时,模型创建被包在
LazyInitContext(default_device=get_accelerator().get_current_device())中,避免在显存中完整物化模型参数,之后再交给 Gemini 进行参数/优化器状态的分区与卸载;其他插件走nullcontext(); - HybridAdam:来自
colossalai.nn,是 ColossalAI 融合实现的自适应优化器,与 fp16 缩放(initial_scale=2**5)配套。
训练循环(train.py)中,前向、optimizer.backward(loss)、梯度裁剪clip_grad_norm_(..., 0.5)与optimizer.step()的耗时被分段计时,并据此计算TFLOPS(公式为model_numel * batch_size * seq_len * 8 / 1e12 / step_time),在 warmup 之后收集各步吞吐并输出中位数,作为插件的性能参考。
与之相对的--distplan=pytorch分支则是一份"朴素对比"实现:num_tokens=256, dim=512, depth=8的小模型 +torch.optim.Adam+ 普通梯度累积,可直观对照两种训练路径。
4.4 运行脚本与多卡启动
原 README 给出的最小运行方式是直接执行:
$ python train.py此时train.py通过colossalai.launch_from_torch()从外部启动器(如colossalai run)获得分布式环境信息。仓库内还提供了两个可直接使用的脚本:
test_ci.sh——CI 冒烟脚本,固定
--dummy_data=True,在GPUNUM=1与GPUNUM=4、batch size 为 2 的情况下以--plugin='gemini'走一遍训练并tee run.log:env OMP_NUM_THREADS=12 colossalai run --nproc_per_node ${GPUNUM} --master_port 29505 \ train.py --dummy_data=True --batch_size=${BATCH_SIZE} --plugin='gemini' 2>&1 | tee run.logrun.sh——保留了更早期调用形态的参数封装(
DISTPAN、PLACEMENT、USE_SHARD_INIT等环境变量)。需要注意:与当前 train.py 的 parse_args 相比,run.sh中传递的--tp_degree/--placement/--shardinit等参数已被新的--plugin/--offload_optim_frac/--dummy_data参数体系取代,从源码结构看属于演进过程中的旧脚本;以当前代码为准,应优先使用test_ci.sh的启动参数写法,或直接按 4.1 的参数表自定义命令行。
训练常量集中在 train.py:NUM_BATCHES=10、WARMUP_BATCHES=1、GRADIENT_ACCUMULATE_EVERY=1、LEARNING_RATE=2e-4、VALIDATE_EVERY=100、GENERATE_EVERY=500、GENERATE_LENGTH=512、SEQ_LEN=1024。训练循环默认只跑 10 个 batch;验证与文本生成代码在当前版本中仍处于注释/待完成状态(见 train.py 的 TODO 区块),如有需要可自行按model.generate(inp, GENERATE_LENGTH)的封装(见 2.7)恢复演示。
4.5 环境依赖
仓库内的 requirements.txt 声明:
colossalai >= 0.1.12 torch >= 1.8.1此外train.py的 import 还依赖einops、tqdm等通用库。由于本仓库已将palm_pytorch作为独立 Python 包内置在 examples/language/palm/palm_pytorch 目录中,训练前只需把工作目录置于examples/language/palm下(保证from palm_pytorch import PaLM可解析),并安装好上述依赖与 ColossalAI 本体即可。原 README 中的pip install PaLM-pytorch面向的是独立分发的同名库,读者若在此仓库内运行 train.py,可直接依赖仓库内置包,无需额外安装。
五、学习与排查建议
- 先跑通再换插件:推荐先用
python train.py --dummy_data=True --plugin=torch_ddp验证数据链路与损失下降,再依次切换torch_ddp_fp16、low_level_zero、gemini,最后用test_ci.sh在 4 卡上做吞吐对照; - 关注显存策略差异:
gemini的显存收益主要来自offload_optim_frac(默认 1.0 全量卸载优化器状态)与LazyInitContext的延迟实例化;low_level_zero则依赖initial_scale的 fp16 缩放,两者都需要搭配HybridAdam而非普通torch.optim.Adam(后者不含initial_scale参数); - 理解 TFLOPS 输出:
logger.info会按 rank 0 打印每步 Loss、Step/FWD/BWD/OPTIM 时间与 TFLOPS,并在结束时给出中位数(train.py),是衡量不同插件吞吐差异的第一手指标; - 回归模型本质:若只关心 PaLM 架构本身,可直接阅读 palm_pytorch/palm_pytorch.py,它把并行残差、RoPE、SwiGLU、共享键值注意力等设计浓缩在少量代码中,是理解后续大型 LLM 结构演化的良好起点。
六、小结
本文以 examples/language/palm/README.md 为骨架,逐层展开了 PaLM 教学实现(palm_pytorch)的架构细节,并以 train.py 为主线说明了 Enwik8 数据加载、自回归训练、HybridAdam + Booster 插件装配、吞吐统计以及test_ci.sh/run.sh的多卡启动方式。该示例兼具两层价值:对模型学习者,它用极简代码呈现了现代 decoder-only LM 的核心零件;对分布式训练实践者,它是一份可直接运行、可横向对比 TorchDDP / Gemini / Low-Level Zero 插件行为的参考基线。
【免费下载链接】ColossalAIMaking large AI models cheaper, faster and more accessible项目地址: https://gitcode.com/GitHub_Trending/co/ColossalAI
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考