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 PaLMPathways 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_head64, heads8, ff_mult4)工厂式构建见 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, eps1e-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 等人与 EleutherAIGPT-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, devicedevice) 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, persistentFalse)对掩码与位置编码做按需缓存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, mult4): inner_dim int(dim * mult) return nn.Sequential( LayerNorm(dim), nn.Linear(dim, inner_dim * 2, biasFalse), SwiGLU(), nn.Linear(inner_dim, dim, biasFalse), )前馈部分把输入升维到inner_dim * 2拆成两半一半经F.silu即 SiLU/Swish激活后与另一半逐元素相乘形成 SwiGLU 门控。源码注释点明这是 Noam Shazeer 论文中的经典做法只是本实现选择 SwiGLU 而非更流行的 GEGLU。注意每个前馈模块开头同样先过无偏置 LayerNorm。2.5 AttentionMulti-Query 键值 因果掩码Attention的核心设计有几点self.to_q将输入投影为heads * dim_head维的 Qself.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_head64, heads8, ff_mult4): net nn.Sequential( nn.Embedding(num_tokens, dim), *[ParallelResidual( Attention(dimdim, dim_headdim_head, headsheads), FeedForward(dimdim, multff_mult), ) for _ in range(depth)], LayerNorm(dim), nn.Linear(dim, num_tokens, biasFalse), ) net[-1].weight net[0].weight # embedding 与输出投影共享权重 nn.init.normal_(net[0].weight, std0.02) return net顶层把depth层并行残差块串起来输出 Linear 的权重直接绑定回 Embedding 权重weight tying随后以标准差 0.02 的正态分布初始化词向量。其余默认超参数为dim_head64, heads8, ff_mult4。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_thres0.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)输入为(batch1, seq_len2048)的 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可确认当前支持的命令行参数如下参数类型默认值说明--distplanstrcolossalai分布式方案colossalai走 Booster 新 APIpytorch走原生 PyTorch 训练分支--pluginstrtorch_ddpBooster 插件可选torch_ddp、torch_ddp_fp16、gemini、low_level_zero--offload_optim_fracfloat1.0优化器状态卸载比例仅对 gemini 插件生效--batch_sizeint8每个 DP 组的 batch size--dummy_databoolFalse是否使用随机整数构造的假数据参数校验在模块加载阶段即完成若--distplan不是colossalai/pytorch会直接抛出TypeError。4.2 数据加载Enwik8 与随机假数据数据准备逻辑见 train.py 中 generate_dataset真实数据从./data/enwik8.gz读取data/目录内的 README 说明 enwik8 数据源自 Hutter prize 页面读取前 95e6 字节并按 90e6 / 5e6 切分为训练集与验证集每字节按uint8词表处理假数据--dummy_dataTrue时直接torch.randint(0, 100, ...)生成 9000 万 / 500 万长度的训练/验证序列便于不下载数据即可进行 CI 与冒烟测试。TextSamplerDatasettrain.py每次随机采样一段长度为seq_len 1的连续字节作为一条样本多出的 1 个 token 供自回归标签使用并用cycle(DataLoader(...))无限循环供给训练循环。4.3 ColossalAI 分支模型、优化器与 Booster 装配在--distplancolossalai时train.pymodel PaLM(num_tokens50304, dim4096, depth64) model AutoregressiveWrapper(model, max_seq_lenSEQ_LEN) optimizer HybridAdam(model.parameters(), lrLEARNING_RATE, initial_scale2**5) model, optimizer, _, _, _ booster.boost(model, optimizer)该分支默认构造一个dim4096、depth64的较大模型以验证分布式插件的承载能力。几个值得注意的底层配合插件选择torch_ddp/torch_ddp_fp16使用TorchDDPPlugin()fp16 时额外注入mixed_precisionfp16gemini使用GeminiPlugin(offload_optim_fracargs.offload_optim_frac, initial_scale2**5)low_level_zero使用LowLevelZeroPlugin(initial_scale2**5)LazyInitContext 与 Gemini 的配合当插件为 gemini 时模型创建被包在LazyInitContext(default_deviceget_accelerator().get_current_device())中避免在显存中完整物化模型参数之后再交给 Gemini 进行参数/优化器状态的分区与卸载其他插件走nullcontext()HybridAdam来自colossalai.nn是 ColossalAI 融合实现的自适应优化器与 fp16 缩放initial_scale2**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 之后收集各步吞吐并输出中位数作为插件的性能参考。与之相对的--distplanpytorch分支则是一份朴素对比实现num_tokens256, dim512, depth8的小模型 torch.optim.Adam 普通梯度累积可直观对照两种训练路径。4.4 运行脚本与多卡启动原 README 给出的最小运行方式是直接执行$ python train.py此时train.py通过colossalai.launch_from_torch()从外部启动器如colossalai run获得分布式环境信息。仓库内还提供了两个可直接使用的脚本test_ci.sh——CI 冒烟脚本固定--dummy_dataTrue在GPUNUM1与GPUNUM4、batch size 为 2 的情况下以--plugingemini走一遍训练并tee run.logenv OMP_NUM_THREADS12 colossalai run --nproc_per_node ${GPUNUM} --master_port 29505 \ train.py --dummy_dataTrue --batch_size${BATCH_SIZE} --plugingemini 21 | 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.pyNUM_BATCHES10、WARMUP_BATCHES1、GRADIENT_ACCUMULATE_EVERY1、LEARNING_RATE2e-4、VALIDATE_EVERY100、GENERATE_EVERY500、GENERATE_LENGTH512、SEQ_LEN1024。训练循环默认只跑 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_dataTrue --plugintorch_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),仅供参考
