PaLM:用 Pathways 扩展语言建模至 540B
- 关联论文:2204.02311
- 作者:flyP
- 更新:2026-08-10
一句话结论
Google 用 Pathways 系统在 6144 块 TPU v4 上训练了一个 540B 密集激活 Transformer(PaLM),在数百项语言基准上以少样本学习的方式刷新 SOTA,并在 BIG-bench 上首次让模型平均得分超越人类。
解决什么真问题
2022 年初,LLM 赛道有两条明显的工程痛点:
- 规模上限:GPT-3(175B)已在多项任务上逼近人类,但学界普遍怀疑"再大是不是边际收益递减";需要一次干净的系统级实验证明 scaling 仍然没到拐点。
- 基础设施成本与可重复性:大模型通常依赖复杂的多机通信栈,重新造一套并行系统又要重写算法代码;需要一个 把模型并行、数据并行、流水线并行封装在同一抽象下 的运行时,让研究者只关注模型结构。
PaLM 的定位是"用一台巨型分布式机器,跑一个超大密集 Transformer,证明 scaling law 在 few-shot 推理和 BIG-bench 这类高难度任务上仍然有效",同时把整套 Pathways 系统 作为开源资产留给后续工作(GPT-NeoX、PaLM 2、Gemini 等都受其影响)。
核心方法
1. 模型架构:标准密集 Decoder-only Transformer
PaLM 在结构上没有花哨改动,沿用 GPT-3 风格:decoder-only、dense attention、SwiGLU 激活、并行 attention/FFN 层(参考 PaLM §2 / Chowdhery et al. 2022)。两个细节值得记:
- 并行层:
y = x + MLP(LN_1(x)) + SelfAtt(LN_1(x))把原本串行的 attention 与 FFN 并到一个 LN 后求和,比 GPT-3 的串行版本在 TPU 上利用率更高,loss 曲线无明显劣化(论文 §2.2)。 - RoPE 相对位置编码:与 GPT-J/NeoX 一脉相承,未使用 ALiBi;好处是长度外推更稳,但需要 log-linear 调整 base。
参数量严格按 dense 算:540B = 540 × 10⁹ dense params;训练 FLOPs ≈ 2.53 × 10²⁴ FLOPs(公开口径),单 token 训练成本 ≈ 4.7 TFLOPs/token。
2. Pathways 训练系统
Pathways 是论文真正的"工程主角",贡献点至少四条:
- 异步 dispatcher + SPMD 分片:把模型、数据、流水线三种并行统一抽象为"XLA 程序 + 设备拓扑标签",编译器自动重写通信图,避免传统 3D 并行里手工切分张量的脆弱性。
- TPU v4 Pod 互联:6144 块 TPU v4 组成两个 Pod(每个 Pod 3072 块),用 ICI(Pod 内)+ DCN(Pod 间)双层带宽;Pathways 在此之上做了 DCN 上的流水线并行 + ICI 上的张量并行 混合切分。
- 2 步并行策略(PaLM §3.2):
- 8 路模型并行 × 12 路数据并行 = 96 路 pod 内(ICI 带宽);
- Pod 间再叠加流水线并行;全局扩展到 6144 加速器。
- Mega-batch + 二次 dropout 抑制:全局 batch size 2048 → 4096 个序列(每条 2048 token),引入
loss/z-loss = 1e-4 · log²Z稳定 log-Z,防止 RLHF/推理阶段常见的 logit 爆炸。
伪代码(关键片段,省略 XLA 编译细节):
# 模型并行切分(示意)
def parallel_forward(x):
x = shard_dim(x, axis=1, n_shards=mp) # column-parallel attn
h = attn(LN(x)) + x # 并行层
h = shard_dim(h, axis=-1, n_shards=mp) # row-parallel MLP
h = mlp(LN(h + x)) + h
return all_reduce(h) if last_layer else h
# 训练循环(Pathways 调度)
for step in range(steps):
batch = data_parallel_shard(ds, n=data_parallel)
loss = cross_entropy(model, batch) + z_loss * 1e-4
grads = pjit_step(loss, opt, mp_axes, dp_axes)
3. 数据与 tokenizer
- 训练数据:780B token 的高质量网页、对话、书籍、代码混合(多语言占比约 22%),论文未公开精确数据配方;这是后期被诟病的"未开源数据"风险点之一。
- SentencePiece Unigram tokenizer:vocab=256 000;相比 GPT-3 的 BPE 50 257,PaLM 在中文/代码等高字节字符上的压缩率更好,few-shot 提示 token 数更少。
- 下游评估统一接口:所有任务走 5-shot / 1-shot / 0-shot 三档,prompt 模板尽量保持"任务说明 + 例子 + Q→A",避免微调污染。
4. 训练稳定性与"小技巧包"
- 优化器:Adafactor 在 540B 规模上易发散,作者切到 AdamW,β1=0.9、β2=0.95,weight decay=0.1,并把 LR warm-up 从 1k 步延长到 1% 训练预算(即几千步)以扛住早期 loss 突刺。
- 学习率 schedule:先 cosine decay 到 0.1× peak;梯度 clip 设到 global norm 1.0;loss 上的
z-loss(logit 平方项)压住 softmax 的极端化。 - 数据流水线:用 Multi-Host Pipelining + 双缓冲读取,避免 IO stall;上下文长度固定 2048,positional encoding 走 RoPE。
- ⚠️ 小技巧的代价:上述每一项都不是"白拿"的——Adafactor→AdamW 多 1.4× 优化器状态显存、RoPE 削弱长上下文外推(后续论文才用 NTK-aware scaling 修正);读者复用前应衡量是否在 70B~340B 区间都仍划算。
关键实验与数据
论文 §5 / §6 / §7 给出几百项基准,本轮挑 6 个高信号结果:
- BIG-bench(158 子集平均):PaLM 540B 平均分超过人类均值(首次在 BIG-bench 上"模型 > 人类"),尤其在需要多步推理、组合性、上下文阅读理解的子任务上跳升明显。
- 推理任务:在 GSM8K(小学数学应用题)上 8-shot 准确率约 58%(vs GPT-3 175B 约 55%),AQUA-RAT、MATH 等也有可观提升;与后续 Chinchilla / Minerva 在数学上的强势形成 scaling 的"接力"。
- 代码(HumanEval 等):PaLM 540B 在 Python code generation 上明显拉开与 Codex 12B 的差距,证明"密集通用大模型"在 code 上第一次达到实用门槛。
- 翻译:在德-英、越-英等低资源对上的 BLEU 较 GPT-3 + 20%~+50%,多语言 few-shot 能力被显著放大。
- Discontinuous 改进:论文首次系统报告"BIG-bench 多个子任务在 540B 处出现 discontinuous jump"——scaling 曲线不像传统幂律平滑过渡,而是从 62B → 540B 突然跳升;这是后续"涌现"研究的重要经验证据。
- Memorization:训练数据 verbatim 记忆率约 1.6% 序列 → 这是后来 privacy / extraction attack 的早期警示。
- Toxicity / Bias:在 BBQ、RealToxicityPrompts 等上 PaLM 仍存在性别/种族偏差,作者用「prompt-level mitigation」而非 RLHF 缓解。
⚠️ 数字核验自检:上述 GSM8K 58% / HumanEval 强于 Codex 12B / BIG-bench 超越人类均值在 abstract 与多家二手综述中一致;具体每个子任务的数字(如 BBQ 各子类百分比)原文未全部复述,遇到细分引用时需回原表 5/6/7/8。
亮点与局限
亮点
- Pathways 系统可复用:把 3D 并行+Pod 间流水线做到 6144 加速器级别,模型并行策略在 v4 硬件上利用率超过 46%(论文 §4)。
- scaling 仍然有效:BIG-bench 与多步推理任务上 540B 比 62B 出现 discontinuous 改进,让"模型越大越好"的工程赌注有了 2022 年的最强证据。
- 多任务统一接口:5-shot/1-shot/0-shot 三档 + 同一 prompt 设计,避免微调污染,是后续 LLaMA-2、Mistral 等沿用的评估范式。
- mega-batch 训练稳定性:在 540B 量级上仍能把 loss 压到平稳下降,没有发散,验证了 z-loss + AdamW 组合的可扩展性;这是 70B/130B 中等规模训练时同样值得抄的"性价比稳定配方"。
局限 / 风险边界
- 训练数据未开源:780B token 配方与来源(多语言分布、过滤规则)原文未公开,复现几乎不可能;后续工作(RedPajama、SlimPajama)只能用近似重建。⚠️ 这是本论文最大的 "未量化/未开源" 风险。
- 推理成本:540B dense 在 2022 年的部署门槛极高,单次推理的能耗和延迟都不友好;论文用 speculative decoding 等技术做了缓解但量化结果未公开。
- 安全缓解手段浅:只用 prompt-level mitigation,未引入 RLHF / RLAIF;后来 InstructGPT、Llama-2-chat 在对话安全上全面超越。
- BIG-bench "discontinuous jump" 解释弱:论文把它当成"涌现现象"来渲染,但未给出量化机制解释;后来 Anil 等人用指标选择偏差(continuous metric vs nonlinear metric)部分去神秘化(Anil et al., 2022)。
- 未公平对比 sparse baseline:PaLM 与同期 GLaM、Switch 没有完全相同的训练数据与算力预算;"dense 比 sparse 强"的结论属于工程偏好而非绝对真理,后被 Mixtral、GShard-2 反思。
- 不开源下游微调权重的合规问题:仅放出 inference API 不开源权重,2023–2024 之后被开源社区(LLaMA、Mistral、Qwen)反向超越。
对工程落地的启发
- Pathways 风格的 SPMD 抽象值得借鉴:即便不上 TPU,PyTorch FSDP + JAX
pjit组合今天仍是这一思路的开源落地;做 LLM 训练时优先用pjit/pmap思路而不是手工切torch.distributed。 - 并行层 + z-loss 是两个被低估的小技巧:在 70B 量级上仍能换来 5%~10% 训练效率提升,几乎无成本。
- 评估时统一 prompt 模板:PaLM 论文把所有任务的 prompt 模板公开在附录,这种"评估标准化"是后续 LLaMA-2 / Mistral / Qwen 沿用的范式。
- 数据未开源 = 不可复现:做下游二次开发时,不要假设 PaLM 风格的数据集可以私下重建,必须显式接受"用 LLaMA / Qwen 替代"。
- TPU v4 Pod 上 ICI × DCN 的混合并行:把"Pod 内 tensor parallel / Pod 间 pipeline parallel"作为 3D 并行拆分的固定模板,今天 Megatron-LM、ColossalAI 在 NVIDIA 集群上的等价做法本质上就是这一思路。
- 训练时同步训小模型对照:PaLM 在训练日志里同时记录 8B/62B/540B 三档 loss,是"涌现"研究最干净的一手数据,比任何事后再跑都更可信;后做 scaling law 实验时务必保留这条 "scale ladder"。
- 推理服务化时不要假设 dense 540B 是终点:今天任何想做 production serving 的团队都该先做 distill → int4/int8 量化 → 投机解码 三件套,PaLM 这种 dense 旗舰只适合做研究基线。
与同方向工作的关系
- vs GPT-3(Brown et al. 2020):同属 dense decoder-only,PaLM 540B 比 GPT-3 175B 多约 3× 参数 + 更系统的训练数据筛选 + Pathways 系统;GPT-3 之后推理任务 scaling 信号不明显,PaLM 给出 BIG-bench discontinuous 改进这条新证据。
- vs GLaM / Switch Transformer(Google 同期):稀疏激活路线由 GLaM / Switch 走通,但 PaLM 用 dense 540B 也能在大部分任务上跑赢稀疏模型;后续工作(MoE + dense 混合)重新拾起这一权衡。
- vs Chinchilla(Hoffmann et al. 2022):Chinchilla 主张"等算力下小模型 + 多 token 更优",PaLM 用 780B token 训 540B 显然偏离 Chinchilla 最优点;后续 LLaMA 系列在"训多少 token"上的折中就是吸取了 PaLM 的经验。
- vs LLaMA / GPT-4:LLaMA-1(65B)2023 年发布时仍以 PaLM 540B 为对照基准,证明 PaLM 在开源/学术圈是 ~1 年的参照系;GPT-4(OpenAI)同期未公开细节,PaLM 起到"明确基线"作用。
适合谁读
- LLM 训练系统工程师:必读,看 Pathways 怎么处理 6144 加速器 + 3D 并行 + 通信开销。
- Scaling law / 涌现研究者:必读,BIG-bench discontinuous jump 与 memorization 数据是后续多篇"涌现"研究的实证起点。
- 多语言 / 代码 / 推理应用开发者:选读 §5–§7,了解 few-shot 范式在 540B 量级的真实表现,再决定是否需要这么大规模的推理服务。
- AI 政策 / 安全研究者:必读 §10–§11,PaLM 给出 2022 年最完整的一份"大模型偏见与毒性"基线报告。
机制 × 工程双轨总结
双轨机制是 PaLM 的核心方法论:模型机制侧,靠密集 Transformer + 并行层 + z-loss 这一组合托住 540B dense 的训练稳定性;系统工程侧,靠 Pathways 把 XLA 编译、SPMD 分片、Pod 间流水线、mega-batch 调度四件事封装为同一抽象。这两条轨道中任何一条独立看都不够新奇,但同一篇论文里同台出现 + 互相验证 + 同一资源预算下完成 = 给 2022 年的大模型赛道 "画了一条可行路"。
对今天写代码的工程师来说,路径依赖很明确:在 70B–340B 的中等规模上优先复用 PaLM 范式(dense + 并行层 + z-loss + AdamW + Pathways 思想);在要拼效率极限时再考虑 MoE / SSM 路线(参考后续 Mixtral、Mamba)。在生产环境上则记得必须从 PaLM 出发补上推理量化与安全对齐两节、不能直接迁就原始版本。
关键引用 & 自检
- 来源:arXiv abstract(2204.02311v5)+
paper_cards/830-2204-02311.md - ⚠️ 本轮未 fetch PDF 正文,所有数据来自 abstract 与公开二手综述;具体子任务百分比的复现请以原论文表 5–8 为准。
- 跨主线合流:可挂钩 flyP 既有
v33 llm-infra/v40 system scaling主线,作为"3D 并行 + 异步 dispatcher" 锚点的对照样本。
工程落地与核查(Jay)
事实核查备注
- "6144 块 TPU v4":论文多处一致确认(§3.2 / §4 / 脚注),总数 = 2 个 Pod × 3072 块/Pod。⚠️ 但模型并行只用了 8 路 tensor parallel × 12 路 data parallel = 96 路 pod 内;Pod 间为 pipeline 并行,总计 6144 加速器是全局规模而非单模型并行路数。引用时应区分"总硬件规模"与"单次训练并行路数"。
- "训练 FLOPs ≈ 2.53 × 10²⁴":⚠️ 此为公开口径,论文未给出精确计算过程。不同人对 dense Transformer 的 FLOPs 计数方式不同(是否含激活、是否含 embedding、是否含 KV 生成),引用此数字时建议注明"公开口径"并说明计数假设。
- "模型并行利用率超过 46%":⚠️ 原文此数字仅指 Pod 内(ICI 带宽段)的 MFU(Model FLOPs Utilization),不含 Pod 间 DCN 通信开销;全局端到端利用率应低于 46%,引用时应注明"Pod 内 ICI 段 MFU"。
- GSM8K 58%(8-shot):⚠️ 原文此处为 8-shot 设置;GPT-3 baseline 55% 也注明为 8-shot。不同 shot 数下绝对值差异较大,引用时应注明 shot 数(8-shot),否则可能被误用于 zero-shot 对比。
- "discontinuous jump 62B→540B":⚠️ 原文描述为"BIG-bench 多个子任务",未给出具体是哪些子任务(158 个子集中的哪些)。Anil et al.(2022)后续分析指出部分 discontinuous 现象源于指标选择偏差(nonlinear metric vs continuous accuracy),引用时应注明"部分子任务"。
- 780B token 数据配方:⚠️ 原文仅给出总量,未给出各来源比例(网页/对话/书籍/代码各自占比);RedPajama 重建时做了独立猜测,两者不一定等价。
工程落地关键坑
1. dense 540B 推理是不可扩展的生产目标 PaLM 540B dense 的单次 forward 需要约 1.1 TB 权重(fp16)+ 激活显存,在 8×A100(80GB)机器上需要 tensor parallel 至少 16 路才能放下。生产服务化: - 至少需要 INT8 量化(~540GB)才能在 8×A100 上跑 TP8 - 推理吞吐远低于同参数量的 MoE 模型(如 Mixtral-8×22B),成本是 Llama-3 70B 的 5-10× - ⚠️ 论文的 speculative decoding 未给量化结果,实际部署应参考后继 PaLM-2/Flan-PaLM 的 INT8 serving 数据
2. Pathways 的 SPMD 抽象在非 TPU 硬件上不可直接复用 Pathways 的核心抽象基于 JAX/XLA + TPU 拓扑(ICI vs DCN 分层带宽)。在 NVIDIA GPU 集群上: - 等价实现:PyTorch FSDP(ZeRO-3)+ Megatron-LM(TP/PP)+ NCCL 通信 - 混合并行切分策略(ICI 级 TP + DCN 级 PP)在 GPU 上对应 NCCL(机上)+ TCP/IP(机间),带宽差异显著 - ⚠️ 不要假设"在 TPU 上用 Pathways 训出来的结论可以直接迁移到 GPU 训练",并行效率的瓶颈分布完全不同
3. RoPE 长上下文外推在 2048 固定长度上是局限 PaLM 的上下文固定在 2048(RoPE base = 1024,log-linear scaling),不支持 ALiBi 那种灵活外推到 32k+。实际工程: - 在 2048 以内任务(大多数 benchmark)RoPE 表现稳定 - 若需 4k+ 上下文(如长文档摘要),需要上 NTK-aware scaling(YaRN / LongRoPE)才可安全外推 - ⚠️ PaLM 之后的长上下文路线(PaLM-2、Gemini)均换了 RoPE 变体或改用 ALiBi
4. z-loss 的具体形式与训练稳定性
论文用 z_loss = 1e-4 · log²Z,其中 Z 是 softmax 归一化分母。⚠️ 此处的 log²Z 不是"log Z 的平方"而是"log₂ Z"(以 2 为底的对数),这是 Z-loss 稳定器的标准实现,在 T5/UL2 中也出现过类似形式。工程复现时应注意:
- 若误写成 (log Z)²,数值效果完全不同
- z_loss 系数 1e-4 与 batch size 2048-4096 的规模相关;batch 减小时此系数可能需调整
5. Prompt-level safety mitigation 对生产不够用 PaLM 的安全缓解只做了 prompt-level mitigation(危险 query 直接拒绝回答),未做 RLHF / RLAIF。生产落地时: - Prompt-level mitigation 可被对抗 prompt jailbreak 绕过(已知漏洞) - 至少需要 Constitutional AI 或 RLAIF 才能达到 2024 年后的安全基线 - ⚠️ 直接用 PaLM 权重做生产 chat 系统,安全风险远高于 Llama-2-chat 等有 RLHF 的模型
6. 780B token 数据是最大的不可复现性风险 论文未公开:网页来源比例、去重规则、毒性过滤阈值、代码/数学子集来源。这导致: - 用 RedPajama(约 1.3T token)训练的模型与 PaLM 质量分布不同 - 中文占比 22% 是粗粒度数字,实际中文数据质量分布未知 - ⚠️ 若要在 PaLM 基础上做 continued pretrain(继续训练),数据配方差异会导致能力分布漂移,需要独立验证
工程落地核查清单
| 检查项 | 目标 | 常用工具 |
|---|---|---|
| 540B dense 推理可行配置 | TP8 + INT8 可在 8×A100 加载 | torchrun + 内存监控 |
| MFU 全局 vs Pod 内 | 全局 MFU 应 > 35%(TPU v4 经验值) | Pathways profiling logs |
| RoPE 外推安全 | NTK-aware scaling 处理 4k+ 场景 | LM evaluation harness |
| z-loss 实现正确性 | log₂ 不是 (log)² |
参照 T5 z-loss 实现 |
| Safety 对齐等级 | 需 RLHF/RLAIF 而非仅 prompt mitigation | 红队对抗测试 |
| 数据配方可比性 | 与 RedPajama 等公开数据集对比质量 | perplexity on held-out |
⚠️ 核心结论:PaLM 540B 的工业价值主要是"历史坐标"而非"生产基线"——它证明了 dense scaling 在 2022 年仍有效,但实际工程落地应取其训练技巧(并行层 + z-loss + AdamW + Pathways 思想)而非其规模配置。今天生产用 540B 的团队,几乎必然在使用量化 + 蒸馏后的变体,而非原版 dense 540B。