Files
2026-07-13 13:37:02 +08:00

1430 lines
104 KiB
HTML
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>Distributed Training 面试 Cheat Sheet</title>
<meta name="generator" content="ARIS render-html (academic, v1)">
<meta name="aris:source-path" content="docs/tutorials/distributed_training_tutorial.md">
<meta name="aris:source-sha256" content="86087779b70e3bcaf4b22536bdf96eae5fd0d9e07675f4b9c61122ce67ee0af7">
<meta name="aris:generated-at" content="2026-05-19 05:38 UTC">
<!-- MathJax 3 -->
<script>
window.MathJax = {
tex: { inlineMath: [['$', '$'], ['\\(', '\\)']], displayMath: [['$$', '$$'], ['\\[', '\\]']], processEscapes: true },
options: { skipHtmlTags: ['script', 'noscript', 'style', 'textarea', 'pre', 'code'] }
};
</script>
<script src="https://cdn.jsdelivr.net/npm/mathjax@3/es5/tex-mml-chtml.js" async></script>
<!-- highlight.js -->
<link rel="stylesheet" href="https://cdn.jsdelivr.net/gh/highlightjs/cdn-release@11.9.0/build/styles/atom-one-light.min.css">
<script src="https://cdn.jsdelivr.net/gh/highlightjs/cdn-release@11.9.0/build/highlight.min.js"></script>
<script>document.addEventListener('DOMContentLoaded', () => hljs.highlightAll());</script>
<style>
:root {
--bg: #fdfcf7;
--bg-soft: #f4f1ea;
--bg-code: #f8f5ec;
--ink: #1a1a1a;
--ink-soft: #4a4a4a;
--ink-muted: #6b6b6b;
--primary: #1a4a8c;
--primary-soft: #2d6cb8;
--accent: #b8390e;
--warn: #b45309;
--warn-bg: #fef3c7;
--info-bg: #dbeafe;
--good-bg: #d1fae5;
--good: #065f46;
--bad-bg: #fee2e2;
--bad: #991b1b;
--border: #d6d0c0;
--border-soft: #e8e3d5;
}
* { box-sizing: border-box; }
html { scroll-behavior: smooth; }
body {
font-family: "Source Serif Pro", "Source Serif 4", "Crimson Pro", "Georgia", "Songti SC", "STSong", serif;
line-height: 1.65;
color: var(--ink);
background: var(--bg);
margin: 0;
padding: 0;
font-size: 16px;
}
.layout {
max-width: 1280px;
margin: 0 auto;
display: grid;
grid-template-columns: 260px 1fr;
gap: 48px;
padding: 40px 32px;
}
nav.toc {
position: sticky;
top: 24px;
align-self: start;
font-size: 13px;
max-height: calc(100vh - 48px);
overflow-y: auto;
border-right: 1px solid var(--border-soft);
padding-right: 16px;
}
nav.toc h3 {
margin: 0 0 12px;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--ink-muted);
font-weight: 600;
}
nav.toc ol { list-style: none; padding: 0; margin: 0; counter-reset: toc; }
nav.toc ol li { margin: 5px 0; counter-increment: toc; }
nav.toc ol li::before { content: counter(toc) ". "; color: var(--ink-muted); margin-right: 4px; }
nav.toc a {
color: var(--ink-soft);
text-decoration: none;
border-bottom: 1px dotted transparent;
}
nav.toc a:hover { color: var(--primary); border-bottom-color: var(--primary); }
nav.toc ul { list-style: none; padding-left: 14px; margin: 3px 0; font-size: 12px; }
nav.toc ul li::before { content: "→ "; color: var(--border); }
main { min-width: 0; }
header.hero {
border-bottom: 3px double var(--primary);
padding-bottom: 24px;
margin-bottom: 32px;
}
header.hero .eyebrow {
color: var(--accent);
font-size: 13px;
text-transform: uppercase;
letter-spacing: 0.12em;
font-weight: 600;
margin-bottom: 8px;
}
header.hero h1 {
font-size: 32px;
line-height: 1.2;
margin: 0 0 12px;
color: var(--ink);
font-weight: 700;
letter-spacing: -0.01em;
}
header.hero .subtitle {
font-size: 16px;
color: var(--ink-soft);
margin: 0 0 8px;
font-style: italic;
}
header.hero .byline {
font-size: 14px;
color: var(--ink-soft);
margin: 0 0 20px;
}
header.hero .byline strong {
color: var(--ink);
font-weight: 600;
}
header.hero .meta {
display: flex;
gap: 20px;
flex-wrap: wrap;
font-size: 12px;
color: var(--ink-muted);
border-top: 1px solid var(--border-soft);
padding-top: 14px;
}
header.hero .meta span strong { color: var(--ink-soft); }
header.hero .meta code {
font-family: "JetBrains Mono", "SF Mono", "Menlo", "Consolas", monospace;
font-size: 11px;
background: var(--bg-soft);
padding: 1px 5px;
border-radius: 3px;
border: 1px solid var(--border-soft);
}
h2 {
font-size: 24px;
margin: 44px 0 14px;
padding-bottom: 8px;
border-bottom: 1px solid var(--border);
color: var(--ink);
font-weight: 700;
}
h2 .num { color: var(--primary); font-weight: 600; margin-right: 8px; }
h3 { font-size: 19px; margin: 28px 0 10px; color: var(--primary); font-weight: 600; }
h4 { font-size: 16px; margin: 20px 0 8px; color: var(--ink); font-weight: 600; }
p { margin: 10px 0; }
ul, ol { padding-left: 22px; margin: 10px 0; }
ul li, ol li { margin: 4px 0; }
ul li::marker { color: var(--primary); }
strong { color: var(--accent); font-weight: 600; }
em { color: var(--ink-soft); }
a { color: var(--primary); }
a:hover { color: var(--accent); }
code:not(.hljs) {
font-family: "JetBrains Mono", "SF Mono", "Menlo", "Consolas", monospace;
font-size: 0.86em;
background: var(--bg-code);
padding: 1px 5px;
border-radius: 3px;
border: 1px solid var(--border-soft);
color: var(--accent);
}
pre {
background: #fafaf6;
border: 1px solid var(--border);
border-left: 4px solid var(--primary);
padding: 0;
overflow-x: auto;
border-radius: 4px;
margin: 14px 0;
}
pre code, pre code.hljs {
background: transparent !important;
display: block;
padding: 14px 18px !important;
font-size: 13px;
line-height: 1.55;
font-family: "JetBrains Mono", "SF Mono", "Menlo", monospace;
color: var(--ink);
}
pre.diagram {
background: #f9f6ed;
border-left: 4px solid var(--accent);
font-size: 12.5px;
line-height: 1.4;
}
.callout {
margin: 16px 0;
padding: 12px 16px;
border-radius: 4px;
border-left: 4px solid;
font-size: 15px;
}
.callout-title {
font-weight: 600;
margin-bottom: 6px;
font-size: 12px;
text-transform: uppercase;
letter-spacing: 0.06em;
}
.callout-info { background: var(--info-bg); border-left-color: var(--primary); }
.callout-info .callout-title { color: var(--primary); }
.callout-warn { background: var(--warn-bg); border-left-color: var(--warn); }
.callout-warn .callout-title { color: var(--warn); }
.callout-good { background: var(--good-bg); border-left-color: var(--good); }
.callout-good .callout-title { color: var(--good); }
.callout-bad { background: var(--bad-bg); border-left-color: var(--bad); }
.callout-bad .callout-title { color: var(--bad); }
table {
width: 100%;
border-collapse: collapse;
margin: 16px 0;
font-size: 14px;
border: 1px solid var(--border);
border-radius: 4px;
overflow: hidden;
}
thead { background: var(--primary); color: white; }
th, td {
text-align: left;
padding: 9px 12px;
border-bottom: 1px solid var(--border-soft);
vertical-align: top;
}
th { font-weight: 600; font-size: 13px; letter-spacing: 0.02em; }
tr:last-child td { border-bottom: none; }
tbody tr:nth-child(even) { background: var(--bg-soft); }
details.qa, details {
background: white;
border: 1px solid var(--border-soft);
border-radius: 6px;
margin: 10px 0;
padding: 0;
}
details summary {
cursor: pointer;
padding: 10px 14px;
font-weight: 600;
font-size: 14px;
color: var(--primary);
list-style: none;
user-select: none;
}
details summary::-webkit-details-marker { display: none; }
details summary::before {
content: "▸ ";
margin-right: 4px;
display: inline-block;
transition: transform 0.15s;
}
details[open] summary::before { transform: rotate(90deg); }
details[open] summary { border-bottom: 1px solid var(--border-soft); }
details > :not(summary) { padding: 10px 14px; }
details p:first-of-type { margin-top: 8px; }
mjx-container[display="true"] { margin: 12px 0 !important; }
footer.aris-footer {
margin-top: 60px;
padding-top: 20px;
border-top: 1px solid var(--border);
font-size: 12px;
color: var(--ink-muted);
}
footer.aris-footer a { color: var(--ink-muted); border-bottom: 1px dotted var(--border); }
@media (max-width: 900px) {
.layout { grid-template-columns: 1fr; gap: 20px; padding: 20px 16px; }
nav.toc {
position: static;
max-height: none;
border-right: none;
border-bottom: 1px solid var(--border-soft);
padding-right: 0;
padding-bottom: 14px;
}
header.hero h1 { font-size: 24px; }
h2 { font-size: 20px; }
}
@media print {
nav.toc { display: none; }
.layout { grid-template-columns: 1fr; padding: 0; }
body { background: white; }
header.hero { border-bottom-color: var(--ink); }
}
</style>
</head>
<body>
<div class="layout">
<nav class="toc">
<h3>Contents</h3>
<ol>
<li><a href="#0-tldr-cheat-sheet">§0 TL;DR Cheat Sheet</a>
</li>
<li><a href="#1-直觉--为什么模型大了一张卡装不下">§1 直觉 — 为什么模型大了一张卡装不下</a>
</li>
<li><a href="#2-nccl-通信原语与拓扑">§2 NCCL 通信原语与拓扑</a>
<ul>
<li><a href="#21-五大集合通信原语">2.1 五大集合通信原语</a></li>
<li><a href="#22-nvlink--ib--拓扑">2.2 NVLink / IB / 拓扑</a></li>
<li><a href="#23-pytorch-中的-nccl-调用">2.3 PyTorch 中的 NCCL 调用</a></li>
</ul>
</li>
<li><a href="#3-ddp--distributeddataparallel">§3 DDP — DistributedDataParallel</a>
<ul>
<li><a href="#31-算法骨架">3.1 算法骨架</a></li>
<li><a href="#32-bucket-fusion--overlapddp-工程精髓">3.2 Bucket Fusion + OverlapDDP 工程精髓)</a></li>
<li><a href="#33-pytorch-代码含-overlap">3.3 PyTorch 代码(含 overlap</a></li>
<li><a href="#34-复杂度">3.4 复杂度</a></li>
</ul>
</li>
<li><a href="#4-zero-123--zero-redundancy-optimizer">§4 ZeRO 1/2/3 — Zero Redundancy Optimizer</a>
<ul>
<li><a href="#41-三阶段显存数学">4.1 三阶段显存数学</a></li>
<li><a href="#42-zero-3-工作流最常用前向--反向--优化">4.2 ZeRO-3 工作流(最常用,前向 / 反向 / 优化)</a></li>
<li><a href="#43-zero-123-vs-ddp-通信对比重要面试题">4.3 ZeRO-1/2/3 vs DDP 通信对比(重要面试题)</a></li>
<li><a href="#44-zero-offload--zero-infinity">4.4 ZeRO-Offload / ZeRO-Infinity</a></li>
</ul>
</li>
<li><a href="#5-fsdp--fsdp2--pytorch-原生-zero-3">§5 FSDP / FSDP2 — PyTorch 原生 ZeRO-3</a>
<ul>
<li><a href="#51-fsdp1-vs-fsdp2-核心区别">5.1 FSDP1 vs FSDP2 核心区别</a></li>
<li><a href="#52-fsdp-wrap-策略最核心的设计决策">5.2 FSDP wrap 策略(最核心的设计决策)</a></li>
<li><a href="#53-fsdp2--混合精度--激活检查点">5.3 FSDP2 + 混合精度 + 激活检查点</a></li>
<li><a href="#54-fsdp-vs-zero-3-通信量差异l3-题">5.4 FSDP vs ZeRO-3 通信量差异(L3 题)</a></li>
<li><a href="#55-hsdp--hybrid-sharded-dp混合分片">5.5 HSDP — Hybrid Sharded DP(混合分片)</a></li>
</ul>
</li>
<li><a href="#6-zero--通信优化2023">§6 ZeRO++ — 通信优化(2023</a>
<ul>
<li><a href="#61-qwzquantized-weight-all-gather量化前向通信">6.1 qwZQuantized Weight all-gather(量化前向通信)</a></li>
<li><a href="#62-hpzhierarchical-partition分层切分">6.2 hpZHierarchical Partition(分层切分)</a></li>
<li><a href="#63-qgzquantized-gradient-reduce量化梯度通信">6.3 qgZQuantized Gradient reduce(量化梯度通信)</a></li>
<li><a href="#64-综合效果">6.4 综合效果</a></li>
</ul>
</li>
<li><a href="#7-tensor-parallel--megatron-lm">§7 Tensor Parallel — Megatron-LM</a>
<ul>
<li><a href="#71-核心思想列并行--行并行-配对">7.1 核心思想:列并行 → 行并行 配对</a></li>
<li><a href="#72-通信量分析">7.2 通信量分析</a></li>
<li><a href="#73-attention-的-tp-切法">7.3 Attention 的 TP 切法</a></li>
<li><a href="#74-tp-代码骨架">7.4 TP 代码骨架</a></li>
<li><a href="#75-tp-必须塞进节点nvlink">7.5 TP 必须塞进节点(NVLink</a></li>
</ul>
</li>
<li><a href="#8-sequence-parallel--省-activation-memory">§8 Sequence Parallel — 省 activation memory</a>
<ul>
<li><a href="#81-动机activation-memory-主要是-layernorm--dropout-这种-element-wise-op">8.1 动机:activation memory 主要是 LayerNorm / Dropout 这种 element-wise op</a></li>
<li><a href="#82-sp-解法沿-sequence-维切">8.2 SP 解法:沿 sequence 维切</a></li>
<li><a href="#83-sp-activation-memory-省了多少l3-题">8.3 SP activation memory 省了多少(L3 题)</a></li>
<li><a href="#84-selective-activation-recompute">8.4 Selective Activation Recompute</a></li>
</ul>
</li>
<li><a href="#9-pipeline-parallel">§9 Pipeline Parallel</a>
<ul>
<li><a href="#91-朴素-pp-与-bubble">9.1 朴素 PP 与 bubble</a></li>
<li><a href="#92-gpipehuang-et-al-neurips-2019">9.2 GPipeHuang et al., NeurIPS 2019</a></li>
<li><a href="#93-1f1b--pipedreamnarayanan-sosp-2019-megatron-lm-2-sc-2021">9.3 1F1B / PipeDreamNarayanan SOSP 2019, Megatron-LM-2 SC 2021</a></li>
<li><a href="#94-interleaved-1f1bmegatron-lm-2-sc-2021">9.4 Interleaved 1F1BMegatron-LM-2, SC 2021</a></li>
<li><a href="#95-1f1b-调度伪代码">9.5 1F1B 调度伪代码</a></li>
<li><a href="#96-pp-通信特点">9.6 PP 通信特点</a></li>
<li><a href="#97-pp-的硬伤load-imbalance">9.7 PP 的硬伤:load imbalance</a></li>
</ul>
</li>
<li><a href="#10-context-parallel--长序列切分">§10 Context Parallel — 长序列切分</a>
<ul>
<li><a href="#101-ring-attentionliu-et-al-arxiv-231001889-2024">10.1 Ring AttentionLiu et al., arXiv 2310.01889, 2024</a></li>
<li><a href="#102-llama-3-cp-实现">10.2 Llama 3 CP 实现</a></li>
<li><a href="#103-cp-与-tp-的正交性">10.3 CP 与 TP 的正交性</a></li>
</ul>
</li>
<li><a href="#11-expert-parallel--moe-路由">§11 Expert Parallel — MoE 路由</a>
<ul>
<li><a href="#111-moe-基本结构">11.1 MoE 基本结构</a></li>
<li><a href="#112-expert-parallelexpert-分到不同-gpu">11.2 Expert Parallelexpert 分到不同 GPU</a></li>
<li><a href="#113-ep-all-to-all-代码骨架伪代码">11.3 EP all-to-all 代码骨架(伪代码)</a></li>
</ul>
</li>
<li><a href="#12-activation-memory-优化">§12 Activation Memory 优化</a>
<ul>
<li><a href="#121-gradient-checkpointingchen-et-al-arxiv160406174-2016">12.1 Gradient CheckpointingChen et al., arXiv:1604.06174, 2016</a></li>
<li><a href="#122-selective-recomputekorthikanti-2023">12.2 Selective RecomputeKorthikanti 2023</a></li>
<li><a href="#123-offloadzero-infinity">12.3 OffloadZeRO-Infinity</a></li>
<li><a href="#124-activation-memory-公式必背">12.4 Activation Memory 公式(必背)</a></li>
</ul>
</li>
<li><a href="#13-综合3d--4d--5d-parallelism">§13 综合:3D / 4D / 5D Parallelism</a>
<ul>
<li><a href="#131-llama-3-405b-训练拓扑公开信息">13.1 Llama 3 405B 训练拓扑(公开信息)</a></li>
<li><a href="#132-deepseek-v3-训练拓扑公开信息">13.2 DeepSeek-V3 训练拓扑(公开信息)</a></li>
</ul>
</li>
<li><a href="#14-dualpipe--2024-流水线前沿">§14 DualPipe — 2024 流水线前沿</a>
<ul>
<li><a href="#141-核心想法">14.1 核心想法</a></li>
<li><a href="#142-dualpipe-vs-普通-1f1b-性质">14.2 DualPipe vs 普通 1F1B 性质</a></li>
<li><a href="#143-dualpipe-的代价">14.3 DualPipe 的代价</a></li>
</ul>
</li>
<li><a href="#15-torchtitan--pytorch-原生-4d-平台">§15 TorchTitan — PyTorch 原生 4D 平台</a>
<ul>
<li><a href="#151-设计目标">15.1 设计目标</a></li>
<li><a href="#152-代码风格与-deepspeed-monkey-patch-对比">15.2 代码风格(与 DeepSpeed monkey patch 对比)</a></li>
<li><a href="#153-float8-训练hopper--blackwell">15.3 Float8 训练(Hopper / Blackwell</a></li>
<li><a href="#154-异步-checkpointing">15.4 异步 Checkpointing</a></li>
</ul>
</li>
<li><a href="#16-通信原语总结表">§16 通信原语总结表</a>
</li>
<li><a href="#17-25-高频面试题">§17 25 高频面试题</a>
<ul>
<li><a href="#l1-必会题任何-ml-工程岗都会问">L1 必会题(任何 ML 工程岗都会问)</a></li>
<li><a href="#l2-进阶题research-oriented-岗位">L2 进阶题(research-oriented 岗位)</a></li>
<li><a href="#l3-高级变体顶级-lab--自研基建岗">L3 高级变体(顶级 lab / 自研基建岗)</a></li>
</ul>
</li>
<li><a href="#a-附录完整-4d-wrap-代码骨架">§A 附录:完整 4D wrap 代码骨架</a>
</li>
<li><a href="#b-参考资料">§B 参考资料</a>
</li>
</ol>
</nav>
<main>
<header class="hero">
<div class="eyebrow">Interview Prep · Distributed Training Systems</div>
<h1>Distributed Training 面试 Cheat Sheet</h1>
<p class="subtitle">DDP / FSDP2 / ZeRO / TP / PP / SP / CP / EP — 公式推导 + From-Scratch 代码 + 25 高频题(L1 必会 · L2 进阶 · L3 顶级 lab)</p>
<p class="byline">By <strong>Ruofeng Yang (杨若峰), Shanghai Jiao Tong University</strong></p>
<div class="meta">
<span><strong>Source:</strong> <code>docs/tutorials/distributed_training_tutorial.md</code></span>
<span><strong>SHA256:</strong> <code>86087779b70e</code></span>
<span><strong>Rendered:</strong> 2026-05-19 05:38 UTC</span>
</div>
</header>
<h2 id="0-tldr-cheat-sheet">§0 TL;DR Cheat Sheet</h2>
<div class="callout callout-info"><div class="callout-title">一页搞定分布式训练 7 大维度</div><p>DP(数据切分) × TP(张量切分) × PP(层切分) × SP(序列切分) × CP(context 切分) × EP(专家切分) × 激活重计算。详见后文 §1–§11 推导。</p></div>
<ol><li><strong>DDP</strong>(PyTorch):每张卡持完整模型,反向时用 NCCL <strong>all-reduce</strong> 同步梯度,bucket fusion + computation/communication overlap 是关键工程优化。</li><li><strong>ZeRO 1/2/3</strong>Rajbhandari et al., SC 2020):把 <strong>optimizer state / gradient / parameter</strong> 分别切到 $N$ 张卡上,单卡显存从 $\Phi (2+2+12) = 16\Phi$ 字节降到 $16\Phi/N$fp16 训练,Adam,参数量 $\Phi$)。</li><li><strong>FSDP / FSDP2</strong>Zhao et al., VLDB 2023PyTorch 2.4+, 2024):PyTorch 原生 ZeRO-3。FSDP2 用 <strong>per-parameter DTensor sharding</strong> 替换 FSDP1 的 flat-parameter 模式,与 TP / PP 复合更顺。</li><li><strong>TP</strong>Megatron-LM, Shoeybi 2019Narayanan SC 2021):<strong>列并行 → 行并行</strong> 配对,每层只需 2× all-reduce/forwardattention 按 head 切,FFN 第一层 col-parallel + 第二层 row-parallel。</li><li><strong>PP</strong>GPipeHuang NeurIPS 2019)→ 1F1B / PipeDreamNarayanan SOSP 2019)→ Megatron-LM <strong>interleaved 1F1B</strong>Narayanan SC 2021)。Bubble ratio $\approx (P-1)/M$$P$ 个 stage$M$ 个 micro-batch),interleaved 把它再压 $V$ 倍。</li><li><strong>SP / CP / EP</strong>SPKorthikanti et al., MLSys 2023)把 TP 外的 LayerNorm/Dropout activation 沿序列切;CPMegatron-LM 2024 / Ring Attention Liu 2024)切长 contextEP 把 MoE expert 分到不同 GPU,前向用 <strong>all-to-all</strong> routing。</li><li><strong>2024 前沿</strong><strong>DualPipe</strong>DeepSeek-V3, Dec 2024)双向流水线把 forward/backward 计算与通信完全 overlap<strong>Llama 3 405B</strong>Meta 2024)用 16K H100、3.8e25 FLOPs、54 天 466 次中断(419 unexpected + 47 计划维护);<strong>TorchTitan</strong>Liang et al., ICLR 2025)打通 FSDP2 + TP + PP + SP + Float8。</li></ol>
<h2 id="1-直觉--为什么模型大了一张卡装不下">§1 直觉 — 为什么模型大了一张卡装不下</h2>
<p><strong>一道送分题但很多人答不全</strong>:训练一个模型,单卡显存里到底装了什么?</p>
<p>考虑 fp16 mixed-precision 训练 + Adam optimizer,参数量为 $\Phi$</p>
<table><thead><tr><th></th><th>单卡占用</th><th>备注</th></tr></thead><tbody><tr><td>Parameter (fp16)</td><td>$2\Phi$</td><td>forward / backward 用</td></tr><tr><td>Gradient (fp16)</td><td>$2\Phi$</td><td>backward 累积</td></tr><tr><td>Optimizer state</td><td>$12\Phi$</td><td>fp32 master copy ($4\Phi$) + Adam $m$, $v$ ($4\Phi + 4\Phi$)</td></tr><tr><td><strong>小计(model states</strong></td><td><strong>$16\Phi$</strong></td><td>与 batch size 无关</td></tr><tr><td>Activation</td><td>$O(B \cdot L \cdot D \cdot \text{depth})$</td><td>与 batch size、seq len 线性,与 depth 线性</td></tr><tr><td>Workspace / temp buffer</td><td>varies</td><td>NCCL / cuDNN workspace</td></tr></tbody></table>
<p>7B 参数模型仅 model states 就要 $\approx 112$ GB——一张 80GB H100 都装不下。所以分布式训练<strong>首要目的是切 model states + activation</strong></p>
<div class="callout callout-info"><div class="callout-title">三个正交切分维度</div><p>训练任务可以沿以下三轴切分,理论上彼此独立、实践中复合使用:</p></div>
<ul><li><strong>数据维度</strong>DP / ZeRO / FSDP):切 batch、模型副本之间通信梯度</li><li><strong>模型维度(深度方向)</strong>PP):切 layerstage 之间传 activation</li><li><strong>模型维度(宽度方向)</strong>TP / SP / CP / EP):切单层内部的张量,每步通信中间结果</li></ul>
<p>3D / 4D / 5D parallelism 就是这几条轴的笛卡尔积。Llama 3 / DeepSeek-V3 都用 4DDP × TP × PP × CP/SP/EP 子集)。</p>
<h2 id="2-nccl-通信原语与拓扑">§2 NCCL 通信原语与拓扑</h2>
<p>分布式训练 99% 的通信交给 NCCL,必须熟悉 5 个原语的语义和通信量。</p>
<h3 id="21-五大集合通信原语">2.1 五大集合通信原语</h3>
<p>设有 $N$ 张 GPU,每张卡持有 size = $S$ 的 buffer。</p>
<table><thead><tr><th>原语</th><th>输入 → 输出</th><th>等价语义</th><th>Ring 算法通信量 / GPU</th></tr></thead><tbody><tr><td><strong>all-reduce</strong></td><td>$N$ 个 $S$ → $N$ 个相同 $S$</td><td>sum + broadcast</td><td>$2(N-1)/N \cdot S \approx 2S$</td></tr><tr><td><strong>reduce-scatter</strong></td><td>$N$ 个 $S$ → $N$ 个不同 $S/N$</td><td>sum 后切片</td><td>$(N-1)/N \cdot S \approx S$</td></tr><tr><td><strong>all-gather</strong></td><td>$N$ 个 $S/N$ → $N$ 个 $S$</td><td>把各 rank 的片段拼成完整</td><td>$(N-1)/N \cdot S \approx S$</td></tr><tr><td><strong>broadcast</strong></td><td>rank0 的 $S$ → $N$ 个 $S$</td><td>单源</td><td>$S$(树形)</td></tr><tr><td><strong>all-to-all</strong></td><td>$N \times N$ 块矩阵转置</td><td>shuffle</td><td>$(N-1)/N \cdot S \approx S$</td></tr></tbody></table>
<p><strong>关键恒等式</strong></p>
<p>$$\boxed{\;\text{all-reduce} = \text{reduce-scatter} + \text{all-gather}\;}$$</p>
<p>每一步 ring 上各传 $S/N$,共 $2(N-1)$ 步 → 单 GPU 总流量 $2(N-1)S/N \approx 2S$ bytes(与 $N$ 几乎无关,这是 ring 算法的精髓)。</p>
<div class="callout callout-warn"><div class="callout-title">面试加分:NCCL 不只一个算法</div><p>NCCL 在小 message 用 <strong>tree all-reduce</strong>latency-bound$O(\log N)$ 跳);大 message 用 <strong>ring all-reduce</strong>bandwidth-bound,吞吐最优)。NVLink 拓扑下还有 <strong>NVLS (NVLink SHARP)</strong>,硬件做 reductionH100/H200 NVSwitch 支持)。</p></div>
<h3 id="22-nvlink--ib--拓扑">2.2 NVLink / IB / 拓扑</h3>
<table><thead><tr><th>链路</th><th>单向带宽(H100 代)</th><th>用途</th></tr></thead><tbody><tr><td>NVLink 4.0</td><td>900 GB/s (per GPU, 18 链路聚合)</td><td>同节点 GPU↔GPU</td></tr><tr><td>PCIe 5.0 x16</td><td>64 GB/s</td><td>GPU↔CPU、慢路径</td></tr><tr><td>InfiniBand NDR 400G</td><td>50 GB/s (per port)</td><td>跨节点</td></tr></tbody></table>
<p><strong>经验法则</strong>:节点内通信比跨节点快 <strong>10-20 倍</strong>。所以 TP 一定要塞进一个节点,DP / PP 才跨节点。Llama 3 训练拓扑:TP=8(节点内 NVLink)× PP=16(跨节点 IB)× DP=128。</p>
<h3 id="23-pytorch-中的-nccl-调用">2.3 PyTorch 中的 NCCL 调用</h3>
<pre><code class="language-python">import torch
import torch.distributed as dist
dist.init_process_group(backend=&quot;nccl&quot;) # 后端固定 nccl
rank = dist.get_rank()
world_size = dist.get_world_size()
# all-reduce(默认 SUM;可选 AVG / MIN / MAX / PRODUCT
buf = torch.ones(1024, device=f&quot;cuda:{rank % 8}&quot;) * rank
dist.all_reduce(buf, op=dist.ReduceOp.SUM)
# 现在 buf == sum(0..world_size-1) * ones(1024)
device = torch.device(f&quot;cuda:{rank % 8}&quot;)
# reduce-scatter
input_list = [torch.full((1024,), float(rank + i), device=device) for i in range(world_size)]
output = torch.empty(1024, device=device)
dist.reduce_scatter(output, input_list, op=dist.ReduceOp.SUM)
# all-gather
local = torch.full((1024,), float(rank), device=device)
gathered = [torch.empty(1024, device=device) for _ in range(world_size)]
dist.all_gather(gathered, local)
# all-to-allMoE 必备)
in_split = list(torch.randn(world_size, 1024, device=device).unbind(0))
out_split = [torch.empty(1024, device=device) for _ in range(world_size)]
dist.all_to_all(out_split, in_split)</code></pre>
<h2 id="3-ddp--distributeddataparallel">§3 DDP — DistributedDataParallel</h2>
<h3 id="31-算法骨架">3.1 算法骨架</h3>
<p>DDP 是数据并行最朴素的实现:</p>
<ol><li><strong>复制</strong>:每个 rank 都持有完整模型副本(参数 / 梯度 / 优化器状态全量)</li><li><strong>切 batch</strong>global batch $B$ 切成 $N$ 份 micro-batch,每 rank 算自己那份的 forward + backward</li><li><strong>同步梯度</strong>backward 完毕后对所有 gradient 做 <strong>all-reduce</strong>(SUM 后除以 $N$ 取平均,等价于 AVG)</li><li><strong>本地 optimizer step</strong>:每 rank 用同样的 gradient 跑同样的 optimizer,参数始终一致</li></ol>
<p>数学形式(loss $\mathcal{L}$ 在 mini-batch 上的均值):</p>
<p>$$g_\text{global} = \frac{1}{N}\sum_{i=1}^N \nabla_\theta \mathcal{L}(\theta; \mathcal{B}_i) = \text{all-reduce-mean}(g_1, \dots, g_N)$$</p>
<p>每个 rank 拿到的 $g_\text{global}$ 完全相同,因此 $N$ 个 rank 上的 $\theta$ 永远保持一致(同初始化 + 同梯度 + 同 optimizer)。</p>
<h3 id="32-bucket-fusion--overlapddp-工程精髓">3.2 Bucket Fusion + OverlapDDP 工程精髓)</h3>
<p>朴素实现:backward 完成 → 把所有梯度按张量拼起来 → 一次 all-reduce。这样 <strong>GPU 空等通信</strong>,浪费惨重。</p>
<p>PyTorch DDP 实际做法:</p>
<ul><li><strong>Bucket</strong>:把多个 gradient tensor 按反向计算顺序打包成 <strong>bucket</strong>(默认 25 MB</li><li><strong>Hook</strong>:在每个 parameter grad 计算完时触发 hook</li><li><strong>重叠</strong>:当一个 bucket 填满(所有 grad 都已计算),<strong>立刻在后台 stream 上发起 all-reduce</strong>,同时主 stream 继续算更早的层的 backward</li><li><strong>结果</strong>:通信被前面层的反向计算完全 / 部分掩盖</li></ul>
<pre class="diagram"><code>Backward time axis →
Layer N: [grad N]──┐
Layer N-1: [grad N-1]┼─bucket_N─[all-reduce N]
Layer N-2: [grad N-2]┘ ↓ (后台 stream)
Layer N-3: [grad N-3]──┐
Layer N-4: [grad N-4]──┼─bucket_N-1─[all-reduce N-1]
... ↓
Layer 1: [grad 1]──────────────[all-reduce 1]
主 stream (compute): ████████████████████████████████
后台 stream (NCCL): ▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓▓
↑ 与计算重叠</code></pre>
<div class="callout callout-warn"><div class="callout-title">bucket 太小 / 太大都不好</div><p>太小:communication launch overhead 主导,bandwidth 利用率低;太大:等齐时间长、首个 all-reduce 启动晚。PyTorch 默认 25 MB 是大模型常用 sweet spot。可通过 <code>DDP(bucket_cap_mb=...)</code> 调。</p></div>
<h3 id="33-pytorch-代码含-overlap">3.3 PyTorch 代码(含 overlap</h3>
<pre><code class="language-python">import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def setup_ddp(rank, world_size, local_rank):
dist.init_process_group(&quot;nccl&quot;, rank=rank, world_size=world_size)
torch.cuda.set_device(local_rank) # 多节点时 local_rank != global rank
def main(rank, world_size, local_rank):
setup_ddp(rank, world_size, local_rank)
device = torch.device(f&quot;cuda:{local_rank}&quot;)
model = MyModel().to(device)
model = DDP(
model,
device_ids=[local_rank],
bucket_cap_mb=25, # bucket 大小
gradient_as_bucket_view=True, # 内存优化: grad 直接是 bucket view
static_graph=False, # 若图静态可开,启用更多 fusion
)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
loader = make_distributed_loader(rank, world_size)
for batch in loader:
x, y = batch[0].to(device), batch[1].to(device)
loss = nn.functional.cross_entropy(model(x), y)
optimizer.zero_grad(set_to_none=True)
loss.backward() # ← bucket hook 在这里触发 all-reduce
optimizer.step()</code></pre>
<h3 id="34-复杂度">3.4 复杂度</h3>
<div class="callout callout-info"><div class="callout-title">记号约定(全文统一)</div><p>后续凡用 $\Phi$ 表示 <strong>参数量</strong>(个数)。fp16 / bf16 训练下一个参数 2 bytes,故 fp16 weight buffer 大小 = $2\Phi$ bytes。下表"$2\Phi$ bytes"指字节数,"$\Phi$ params"指个数。</p></div>
<table><thead><tr><th></th><th></th><th>说明</th></tr></thead><tbody><tr><td>单 GPU 显存</td><td>$16\Phi$ bytes + activation</td><td>model states 不分摊(含 fp16 params/grads + fp32 master/Adam</td></tr><tr><td>单 step 通信量</td><td>$\approx 2 \cdot 2\Phi = 4\Phi$ bytesring all-reduce on fp16 gradient</td><td>等价于 reduce-scatter + all-gather</td></tr><tr><td>扩展性</td><td>计算线性 scale;通信 bandwidth $O(1)$ring 不依赖 $N$</td><td>但 latency-bound 小模型仍受 $N$ 影响</td></tr></tbody></table>
<p><strong>DDP 的硬伤</strong>model states 不切。Adam 训练 70B 模型,单卡光 model states 就 1.12 TBDDP 完全装不下。所以 7B+ 模型必须上 ZeRO / FSDP。</p>
<h2 id="4-zero-123--zero-redundancy-optimizer">§4 ZeRO 1/2/3 — Zero Redundancy Optimizer</h2>
<p>Rajbhandari et al. <strong>"ZeRO: Memory Optimizations Toward Training Trillion Parameter Models"</strong> (SC 2020) 是分布式训练第二个划时代工作(第一个是 Megatron-TP)。核心思想:DDP 把 model states 复制 $N$ 份是浪费,切成 $N$ 份每张卡只放 $1/N$ 即可,通信只多一点点。</p>
<h3 id="41-三阶段显存数学">4.1 三阶段显存数学</h3>
<p>记 $\Phi$ = 参数量,$N$ = DP world size。fp16 mixed precision + Adam 下单卡 model states</p>
<p>$$\text{DDP}: \quad 2\Phi + 2\Phi + 12\Phi = 16\Phi$$</p>
<p>ZeRO 三阶段按 <strong>「切什么」</strong> 区分:</p>
<table><thead><tr><th>阶段</th><th>切的内容</th><th>单卡 model states</th><th>通信量 (per step)</th></tr></thead><tbody><tr><td><strong>ZeRO-1</strong></td><td>optimizer state</td><td>$2\Phi + 2\Phi + 12\Phi/N$</td><td>$2\Phi$all-reduce 等价于 reduce-scatter + all-gatherZeRO-1 只 reduce-scatter grad</td></tr><tr><td><strong>ZeRO-2</strong></td><td>optimizer state + gradient</td><td>$2\Phi + 2\Phi/N + 12\Phi/N$</td><td>$2\Phi$(同 ZeRO-1</td></tr><tr><td><strong>ZeRO-3</strong></td><td>opt state + grad + <strong>parameter</strong></td><td>$2\Phi/N + 2\Phi/N + 12\Phi/N = 16\Phi/N$</td><td>$3\Phi$forward + backward 各一次 all-gatherbackward 一次 reduce-scatter</td></tr></tbody></table>
<div class="callout callout-good"><div class="callout-title">极限情况下</div><p>$N$ 足够大时(如 1024 张 H100),ZeRO-3 单卡 model states $16\Phi / 1024 = 0.0156\Phi$。65B 模型即 $\approx 1$ GB / GPU,配合 activation checkpoint 单 H100 完全装得下。</p></div>
<h3 id="42-zero-3-工作流最常用前向--反向--优化">4.2 ZeRO-3 工作流(最常用,前向 / 反向 / 优化)</h3>
<p>ZeRO-3 把参数本身也切了,所以参数在用之前必须 <strong>all-gather</strong> 才能 forward。</p>
<p><strong>前向流程</strong>(每层 / 每 module):</p>
<pre class="diagram"><code>1. all-gather: 从 N 张卡 collect 完整 W^() ──[通信: φ_ bytes]
2. compute: y = f(x; W^()) ──[计算]
3. release: 丢掉本地不属于自己的 shard ──[释放显存]</code></pre>
<p><strong>反向流程</strong></p>
<pre class="diagram"><code>1. all-gather: 再次取回 W^()forward 时已释放)──[通信: φ_ℓ]
2. compute: grad_W^(), grad_x
3. reduce-scatter: 把 grad_W^() reduce 并切回各自的 shard ──[通信: φ_ℓ]
4. release: 丢掉非己 shard</code></pre>
<p>伪代码:</p>
<pre><code class="language-python"># ZeRO-3 forward (单层抽象, 简化版)
def zero3_forward(layer_idx, x, sharded_W):
# 1. all-gather full weight
full_W = all_gather(sharded_W) # [Φ_ / N, ...] × N -&gt; [Φ_, ...]
# 2. compute
y = layer_forward(x, full_W)
# 3. 释放 full_W (只留 sharded_W)
del full_W
return y
def zero3_backward(layer_idx, dy, sharded_W, cached_input):
# 1. all-gather (forward 时已释放)
full_W = all_gather(sharded_W)
# 2. compute local gradients
dW_local, dx = layer_backward(dy, full_W, cached_input)
del full_W
# 3. reduce-scatter gradient 到对应 shard
dW_sharded = reduce_scatter(dW_local) # [Φ_, ...] / N -&gt; [Φ_/N, ...]
return dW_sharded, dx</code></pre>
<h3 id="43-zero-123-vs-ddp-通信对比重要面试题">4.3 ZeRO-1/2/3 vs DDP 通信对比(重要面试题)</h3>
<p>总参数 $\Phi$(个数)。下表的单位是 <strong>"fp16 weight buffer 等价数"</strong>(即 $\Phi$ 在通信流量列代表 $2\Phi$ bytes 的 fp16 buffer 流量)。一次 forward+backward+update 的 <strong>per-GPU</strong> 通信流量(ring 假设):</p>
<table><thead><tr><th>模式</th><th>Forward</th><th>Backward</th><th>Optim</th><th>总(fp16 buffer equiv.</th></tr></thead><tbody><tr><td>DDP</td><td>0</td><td>$2\Phi$ (all-reduce grad)</td><td>0</td><td>$2\Phi$(即 $4\Phi$ bytes</td></tr><tr><td>ZeRO-1</td><td>0</td><td>$\Phi$ (reduce-scatter grad)</td><td>$\Phi$ (all-gather updated params)</td><td>$2\Phi$</td></tr><tr><td>ZeRO-2</td><td>0</td><td>$\Phi$ (reduce-scatter grad)</td><td>$\Phi$ (all-gather)</td><td>$2\Phi$</td></tr><tr><td>ZeRO-3</td><td>$\Phi$ (all-gather params, on-the-fly)</td><td>$\Phi$ (all-gather) + $\Phi$ (reduce-scatter grad)</td><td>0 (已在 backward 中)</td><td>$3\Phi$</td></tr></tbody></table>
<div class="callout callout-info"><div class="callout-title">关键结论</div><p>ZeRO-1/2 通信量与 DDP 相同($2\Phi$ buffer),但显存大幅下降;ZeRO-3 多 $1.5\times$ 通信,换 $N\times$ 显存下降。实际工程中 ZeRO-3 通信也能通过 <strong>prefetch + overlap</strong> 部分隐藏,所以是 70B+ 模型的主流选择。换算成字节:DDP $\approx 4\Phi$ bytesZeRO-3 $\approx 6\Phi$ bytes。</p></div>
<h3 id="44-zero-offload--zero-infinity">4.4 ZeRO-Offload / ZeRO-Infinity</h3>
<p><strong>ZeRO-Offload</strong>Ren et al., USENIX ATC 2021):把 optimizer state + 部分 gradient 卸到 <strong>CPU</strong>CPU 端跑 Adam update。代价:CPU↔GPU PCIe 通信 + CPU 计算慢。适合<strong>小集群 + 大模型</strong>场景(如单机 8 卡跑 13B)。</p>
<p><strong>ZeRO-Infinity</strong>Rajbhandari et al., SC 2021):再进一步,把参数 / 优化器卸到 <strong>NVMe</strong>。理论上单机能跑 1T 参数(实际 throughput 极低,主要用于推理 / 微调)。</p>
<div class="callout callout-warn"><div class="callout-title">Offload 用不用是 trade-off</div><p>卸 CPU 后单 step 时间通常变长 1.5-3 倍;NVMe 卸更慢 5-10 倍。只在「装不下 + 没钱加卡」时才用。生产训练优先扩 GPU 数。</p></div>
<h2 id="5-fsdp--fsdp2--pytorch-原生-zero-3">§5 FSDP / FSDP2 — PyTorch 原生 ZeRO-3</h2>
<p>Zhao et al. <strong>"PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel"</strong> (VLDB 2023) 把 ZeRO-3 思想做进 PyTorch 主干。FSDP2PyTorch 2.4+, 2024)是大重写:<strong>per-parameter DTensor sharding</strong> 替换 FSDP1 的 flat-parameter。</p>
<h3 id="51-fsdp1-vs-fsdp2-核心区别">5.1 FSDP1 vs FSDP2 核心区别</h3>
<table><thead><tr><th>维度</th><th>FSDP1 (2022-2023)</th><th>FSDP2 (2024+)</th></tr></thead><tbody><tr><td>数据结构</td><td><strong>FlatParameter</strong>:把同一 wrap unit 内所有 param 拼成一个 1D buffer,再 chunk</td><td><strong>DTensor per-parameter</strong>:每个 param 独立按 dim-0 切</td></tr><tr><td>状态字典</td><td>需 all-gather 才能产出 unflattened state dict</td><td>通信自由 sharded state dict</td></tr><tr><td>冻结参数</td><td>同 unit 内必须全冻或全可训</td><td>各 param 独立,混合冻可训自然</td></tr><tr><td>TP 复合</td><td>困难(flat-buffer 与 TP 沿不同维切冲突)</td><td><strong>天然兼容</strong>DTensor 描述多轴 placement (<code>Shard(0)</code>, <code>Replicate</code>, <code>Shard(1)</code> 组合)</td></tr><tr><td>API</td><td><code>FullyShardedDataParallel</code></td><td><code>fully_shard()</code> 函数式 wrap</td></tr></tbody></table>
<h3 id="52-fsdp-wrap-策略最核心的设计决策">5.2 FSDP wrap 策略(最核心的设计决策)</h3>
<p>FSDP 不是把整个模型当一个 unit。它按<strong>自定义 unit boundary</strong> 切,每个 unit 内的参数一起 all-gather / reduce-scatter。</p>
<pre><code class="language-python">import torch
import torch.nn as nn
from torch.distributed.fsdp import fully_shard # FSDP2 API
from torch.distributed.fsdp import MixedPrecisionPolicy
class TransformerBlock(nn.Module):
def __init__(self, d_model, n_heads):
super().__init__()
self.attn = MultiHeadAttention(d_model, n_heads)
self.mlp = FeedForward(d_model)
self.ln1 = nn.LayerNorm(d_model)
self.ln2 = nn.LayerNorm(d_model)
def forward(self, x):
x = x + self.attn(self.ln1(x))
x = x + self.mlp(self.ln2(x))
return x
# 把每个 TransformerBlock 当一个 FSDP unit
model = MyLLM(n_layers=32, d_model=4096, n_heads=32).cuda()
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16, # forward 用 bf16
reduce_dtype=torch.float32, # gradient reduce 用 fp32 防误差累积
)
for block in model.blocks:
fully_shard(block, mp_policy=mp_policy)
fully_shard(model, mp_policy=mp_policy) # root 也要 wrap</code></pre>
<div class="callout callout-warn"><div class="callout-title">wrap 粒度的工程权衡</div><p>Unit 越小(如每个 linear 都 wrap):每次 all-gather 只取一层,显存 peak 低,但通信次数多、prefetch 难做;Unit 越大(如整个 block 或多个 block):通信次数少、容易 overlap,但 peak memory 高。<strong>TransformerBlock 粒度是 LLaMA / GPT 类的标准答案</strong></p></div>
<h3 id="53-fsdp2--混合精度--激活检查点">5.3 FSDP2 + 混合精度 + 激活检查点</h3>
<pre><code class="language-python">from torch.distributed.fsdp import CPUOffloadPolicy, OffloadPolicy
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
checkpoint_wrapper, CheckpointImpl
)
# 1. 激活检查点 (gradient checkpointing) — 重算换显存
def apply_ac(model):
for i, block in enumerate(model.blocks):
model.blocks[i] = checkpoint_wrapper(
block,
checkpoint_impl=CheckpointImpl.NO_REENTRANT,
)
apply_ac(model)
# 2. FSDP2 wrap
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
)
offload_policy = OffloadPolicy() # 或 CPUOffloadPolicy() 卸到 CPU
for block in model.blocks:
fully_shard(block, mp_policy=mp_policy, offload_policy=offload_policy)
fully_shard(model, mp_policy=mp_policy)
# 3. 训练 (注意: 必须用 torch._foreach_ 友好的 optimizer)
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, fused=True)
for batch in loader:
loss = compute_loss(model, batch)
loss.backward()
optimizer.step()
optimizer.zero_grad(set_to_none=True)</code></pre>
<h3 id="54-fsdp-vs-zero-3-通信量差异l3-题">5.4 FSDP vs ZeRO-3 通信量差异(L3 题)</h3>
<p><strong>结论</strong><strong>完全一样</strong>(在 wrap unit = layer 这个粒度上)。FSDP 是 ZeRO-3 的 PyTorch 原生实现,算法等价。差异在工程:</p>
<table><thead><tr><th>维度</th><th>DeepSpeed ZeRO-3</th><th>FSDP2</th></tr></thead><tbody><tr><td>集成</td><td>外部库 + monkey patch</td><td>PyTorch 主干</td></tr><tr><td>state_dict</td><td>需要单独 API 拼装</td><td>原生 <code>set_model_state_dict</code> / DCP</td></tr><tr><td>TP 复合</td><td>DeepSpeed-Megatron 桥接</td><td>DTensor 原生</td></tr><tr><td>冻结 / LoRA</td><td>复杂 (flat-param 冲突)</td><td>自然</td></tr><tr><td>编译</td><td>受限</td><td><code>torch.compile(fullgraph=False)</code> 友好</td></tr></tbody></table>
<p>实际选型:<strong>新项目用 FSDP2DeepSpeed 仍在 70B+ MoE / 多框架场景用得多</strong></p>
<h3 id="55-hsdp--hybrid-sharded-dp混合分片">5.5 HSDP — Hybrid Sharded DP(混合分片)</h3>
<p><strong>问题</strong>1024 张 GPU 全部 ZeRO-3,单次 all-gather 跨 1024 卡 → 单卡通信量 $(N-1)/N \cdot \phi \to \phi$ 不变,但 <strong>latency 大幅上升</strong>IB 跨节点贵)。</p>
<p><strong>HSDP</strong>:组内 ZeRO-3(如节点内 8 卡),组间 DDP。等价于把 1024 卡分成 128 个 ZeRO-3 group(每组 8 卡)。</p>
<table><thead><tr><th>模式</th><th>Shard 范围</th><th>单卡 model states</th><th>跨节点通信</th></tr></thead><tbody><tr><td>纯 ZeRO-3</td><td>全 1024 卡</td><td>$16\Phi / 1024$</td><td></td></tr><tr><td>HSDP (8)</td><td>8 卡组内</td><td>$16\Phi / 8 = 2\Phi$</td><td>少(组间是 grad all-reduce</td></tr></tbody></table>
<p>trade-off<strong>显存换通信效率</strong>。Llama 3 在某些训练阶段用 HSDP。</p>
<h2 id="6-zero--通信优化2023">§6 ZeRO++ — 通信优化(2023</h2>
<p>Wang et al. <strong>"ZeRO++: Extremely Efficient Collective Communication for Giant Model Training"</strong> (NeurIPS MLSys workshop 2023, arXiv 2306.10209)。在 ZeRO-3 之上做了 3 件事:</p>
<h3 id="61-qwzquantized-weight-all-gather量化前向通信">6.1 qwZQuantized Weight all-gather(量化前向通信)</h3>
<p>ZeRO-3 forward 时每层都要 all-gather fp16 weight。<code>qwZ</code> 把 fp16 → int8 后再 all-gather<strong>通信量减半</strong></p>
<pre class="diagram"><code>forward (原始): weight_shard (fp16) ──all-gather──&gt; full_weight (fp16) ──compute
forward (qwZ): weight_shard (fp16) ──quant──&gt; shard (int8)
──all-gather──&gt; full (int8)
──dequant──&gt; full (fp16) ──compute</code></pre>
<p>代价:每次 all-gather 前后做 block-wise quantization / dequantizationblock size 一般 2048-4096 元素一组,保留 per-block scale)。block-quant 比 per-tensor quant 精度好得多,对训练 loss 影响 &lt; 1%。</p>
<h3 id="62-hpzhierarchical-partition分层切分">6.2 hpZHierarchical Partition(分层切分)</h3>
<p>观察:<strong>backward 时 all-gather 的代价比 forward 大</strong>(更深更靠后的层先反传)。<code>hpZ</code> 把 weight 在 <strong>节点内复制</strong>(全 NVLink),跨节点不切:</p>
<ul><li>节点内:8 张 GPU 各持完整 weighthpZ = 节点内 DDP</li><li>跨节点:weight 切片(依旧 ZeRO-3 模式)</li></ul>
<p>backward 的 weight all-gather 只在节点内做(NVLink,便宜),不跨 IB。代价:每节点显存翻 8 倍——但 model states 本来就被切到 1024 卡上,乘 8 还远小于 DDP。</p>
<h3 id="63-qgzquantized-gradient-reduce量化梯度通信">6.3 qgZQuantized Gradient reduce(量化梯度通信)</h3>
<p>backward 末段 reduce-scatter gradient 也走 int8fp16 → int8 + 量化 reduce)。原始 reduce-scatter 用 SUM 操作不能简单量化(量化 + sum 会累积误差),ZeRO++ 改用 <strong>all-to-all + 反量化 + local sum</strong> 路径绕开。</p>
<h3 id="64-综合效果">6.4 综合效果</h3>
<p>ZeRO++ 论文报告:<strong>384 GPU 规模 throughput 2.16× 提升</strong>,通信量降 4×(fp16→int8 节省 2× × 2 个通信原语)。代价:实现复杂、精度需小心 ablation。</p>
<div class="callout callout-warn"><div class="callout-title">量化通信 ≠ 量化训练</div><p>这里只量化<strong>集合通信途中的 transient buffer</strong>,权重本身的存储和计算仍是 fp16 / bf16。所以 loss 影响通常可忽略;和 fp8 训练(如 Hopper 上的 fp8 GEMM)是两件事。</p></div>
<h2 id="7-tensor-parallel--megatron-lm">§7 Tensor Parallel — Megatron-LM</h2>
<p>Shoeybi et al. <strong>"Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism"</strong> (arXiv 2019) + Narayanan et al. <strong>"Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM"</strong> (SC 2021)。</p>
<h3 id="71-核心思想列并行--行并行-配对">7.1 核心思想:列并行 → 行并行 配对</h3>
<p>考虑一个 MLP 层:$Y = \text{GELU}(X W_1) W_2$$X \in \mathbb{R}^{B \times D}$$W_1 \in \mathbb{R}^{D \times 4D}$$W_2 \in \mathbb{R}^{4D \times D}$。</p>
<p><strong>Column Parallel</strong>(切 $W_1$ 的输出维 $4D$):</p>
<p>$$W_1 = [W_1^{(1)} \mid W_1^{(2)} \mid \cdots \mid W_1^{(T)}], \quad W_1^{(i)} \in \mathbb{R}^{D \times 4D/T}$$</p>
<p>每个 rank $i$ 持有 $W_1^{(i)}$,独立算 $Y_1^{(i)} = X W_1^{(i)} \in \mathbb{R}^{B \times 4D/T}$。<strong>输入 $X$ 全副本,输出沿 column 切</strong></p>
<p><strong>Row Parallel</strong>(切 $W_2$ 的输入维 $4D$):</p>
<p>$$W_2 = \begin{bmatrix} W_2^{(1)} \\ W_2^{(2)} \\ \vdots \\ W_2^{(T)} \end{bmatrix}, \quad W_2^{(i)} \in \mathbb{R}^{4D/T \times D}$$</p>
<p>输入沿 row 切:$Z^{(i)} = \text{GELU}(Y_1^{(i)}) W_2^{(i)} \in \mathbb{R}^{B \times D}$。每个 rank 算自己那部分 partial sum,最后 <strong>all-reduce</strong> 求和:</p>
<p>$$Y = \sum_{i=1}^T Z^{(i)} = \text{all-reduce}(Z^{(1)}, \dots, Z^{(T)})$$</p>
<h3 id="72-通信量分析">7.2 通信量分析</h3>
<p>列并行:输入 $X$ 是全副本,输出沿 column 切 → <strong>forward 无通信</strong>(输出已分散在各 rank)。</p>
<p>行并行:输入已沿 row 切,输出需 all-reduce → <strong>forward 1× all-reduce $(BD)$</strong></p>
<p><strong>col + row 配对</strong>:列并行末端不通信(结果直接喂给行并行),行并行末端 1× all-reduce。整个 MLP block forward <strong>只 1× all-reduce</strong></p>
<p>Backward 镜像:每个 MLP block backward 也是 1× all-reducegradient 流过 col-row 时镜像通信)。</p>
<p><strong>总计</strong>:单个 transformer blockMLP + Attentionattention 也是 col-rowforward + backward = <strong>4× all-reduce, 每次 $BLD$</strong>$L$ = seq len)。</p>
<h3 id="73-attention-的-tp-切法">7.3 Attention 的 TP 切法</h3>
<p>Multi-head attention 的 $W_Q, W_K, W_V \in \mathbb{R}^{D \times D}$。直接<strong>按 head 维切</strong>(每个 head 独立、列并行天然):</p>
<pre><code>H heads, T-way TP:
- 每 rank 持 H/T 个 head 的 W_Q^(i), W_K^(i), W_V^(i)
- 每 rank 独立算自己的 head_h(Q W_Q^h, ...)
- 输出 W_O 切行并行 (row-parallel) → 末端 all-reduce</code></pre>
<div class="callout callout-info"><div class="callout-title">为什么 attention 切 head 而不切 dim</div><p>head 维天然独立(不同 head 不交互),切 head 不需要跨 rank 通信中间结果;切 hidden dim 则需要在 softmax 前后做额外通信。代码也简洁:直接把 head 维分配给 ranks。</p></div>
<h3 id="74-tp-代码骨架">7.4 TP 代码骨架</h3>
<p>下面是教学版 col-parallel / row-parallel,明确写出 forward 和 backward 的通信对(关键:col-parallel forward 不通信,backward all-reduce 输入梯度;row-parallel forward all-reduce 输出,backward 不通信):</p>
<pre><code class="language-python">import math
import torch
import torch.nn as nn
import torch.distributed as dist
class _CopyToTPRegion(torch.autograd.Function):
&quot;&quot;&quot; Identity in forward, all-reduce in backward (col-parallel 入口) &quot;&quot;&quot;
@staticmethod
def forward(ctx, x, tp_group):
ctx.tp_group = tp_group
return x
@staticmethod
def backward(ctx, grad_out):
dist.all_reduce(grad_out, op=dist.ReduceOp.SUM, group=ctx.tp_group)
return grad_out, None
class _ReduceFromTPRegion(torch.autograd.Function):
&quot;&quot;&quot; All-reduce in forward, identity in backward (row-parallel 出口) &quot;&quot;&quot;
@staticmethod
def forward(ctx, x, tp_group):
dist.all_reduce(x, op=dist.ReduceOp.SUM, group=tp_group)
return x
@staticmethod
def backward(ctx, grad_out):
return grad_out, None
class ColumnParallelLinear(nn.Module):
&quot;&quot;&quot; Y = X W, W ∈ R^{in×out}, 切 out 维成 T 份;
forward: input replicated -&gt; 输出 [B, out/T] 沿 last dim 切
backward: input grad 要 all-reduce(多个 TP rank 各算了 partial dX&quot;&quot;&quot;
def __init__(self, in_features, out_features, tp_group, tp_size):
super().__init__()
assert out_features % tp_size == 0
self.tp_group = tp_group
self.out_per_rank = out_features // tp_size
self.weight = nn.Parameter(torch.empty(self.out_per_rank, in_features))
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
def forward(self, x):
# x: [B, in] (replicated). 入口 _CopyToTPRegion 让 backward 自动 all-reduce dX
x = _CopyToTPRegion.apply(x, self.tp_group)
return torch.nn.functional.linear(x, self.weight) # [B, out/T]
class RowParallelLinear(nn.Module):
&quot;&quot;&quot; Y = X W, W ∈ R^{in×out}, 切 in 维成 T 份;
forward: input [B, in/T] (sharded), 输出 partial sum -&gt; all-reduce
backward: dW/dX 各自本地算即可,不需通信 &quot;&quot;&quot;
def __init__(self, in_features, out_features, tp_group, tp_size):
super().__init__()
assert in_features % tp_size == 0
self.tp_group = tp_group
self.in_per_rank = in_features // tp_size
self.weight = nn.Parameter(torch.empty(out_features, self.in_per_rank))
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
def forward(self, x):
# x: [B, in/T] (sharded along last dim)
local_out = torch.nn.functional.linear(x, self.weight) # [B, out] partial sum
return _ReduceFromTPRegion.apply(local_out, self.tp_group)
class TPTransformerMLP(nn.Module):
&quot;&quot;&quot; Megatron-LM 标准 col+row 配对 &quot;&quot;&quot;
def __init__(self, d_model, tp_group, tp_size):
super().__init__()
d_ff = 4 * d_model
self.fc1 = ColumnParallelLinear(d_model, d_ff, tp_group, tp_size) # 出 [B, L, 4D/T]
self.fc2 = RowParallelLinear(d_ff, d_model, tp_group, tp_size) # 入 [B, L, 4D/T] -&gt; [B, L, D]
def forward(self, x):
# x: [B, L, D] (replicated)
h = torch.nn.functional.gelu(self.fc1(x)) # [B, L, 4D/T]
return self.fc2(h) # [B, L, D] after all-reduce</code></pre>
<div class="callout callout-info"><div class="callout-title">TP 通信对子(必记)</div><p>col-parallel: <code>forward = identity, backward = all-reduce(dX)</code>row-parallel: <code>forward = all-reduce(Y), backward = identity</code>。一个 transformer block (col + row × 2 sets: MLP + attention) 共 4 次 all-reduce2 在 forward、2 在 backward——这就是上文 §7.2 那 4 次的来源。</p></div>
<h3 id="75-tp-必须塞进节点nvlink">7.5 TP 必须塞进节点(NVLink</h3>
<p>每个 transformer block forward + backward = 4× all-reduce of activation tensor 大小 $\approx BLD$。对 7B 模型 $D = 4096$$B \times L = 2048$ → 单次 $\approx 32$ MBper-block 4 次 × 32 layer = <strong>128 次 / step, $\approx 4$ GB / step</strong></p>
<p>NVLink 带宽 900 GB/s 下 $\approx 4.4$ msIB 50 GB/s 下 $\approx 80$ ms。所以 <strong>TP 必须塞进节点内</strong>(NVLink 域)。这就是为什么 LLaMA / Llama 3 / GPT-3 都用 TP=8(一个节点 8 卡)而不是 TP=16。</p>
<h2 id="8-sequence-parallel--省-activation-memory">§8 Sequence Parallel — 省 activation memory</h2>
<p>Korthikanti et al. <strong>"Reducing Activation Recomputation in Large Transformer Models"</strong> (MLSys 2023)。</p>
<h3 id="81-动机activation-memory-主要是-layernorm--dropout-这种-element-wise-op">8.1 动机:activation memory 主要是 LayerNorm / Dropout 这种 element-wise op</h3>
<p>TP 把 attention / MLP 的中间 activation 切了(沿 head / hidden 维),但 <strong>TP 外的部分</strong>LayerNorm 输入输出、Dropout、residual)仍是<strong>全副本</strong>——每张 TP rank 都存一份 $[B, L, D]$ 大小的 activation。</p>
<p>7B 模型 $B \times L \times D = 2048 \times 4096$, fp16 = 16 MB / activation。一层有约 4-6 个这种 full-replica activation。32 层 × 5 → 2.5 GB / GPU 浪费。</p>
<h3 id="82-sp-解法沿-sequence-维切">8.2 SP 解法:沿 sequence 维切</h3>
<p>让 TP 外的 activation 也沿 $L$sequence)维切到 TP rank 上。</p>
<pre><code>TP-only: [B, L, D] [B, L, D]
全副本 全副本
↓ ↑
LayerNorm → split → Attention → all-gather → LayerNorm
(TP)
TP + SP: [B, L/T, D] [B, L/T, D]
沿 L 切 沿 L 切
↓ ↑
LayerNorm → all-gather → Attention → reduce-scatter → LayerNorm
(TP)</code></pre>
<p><strong>关键操作变化</strong></p>
<ul><li>TP-only: forward 内 attention/MLP 入口处需要 broadcast / 退化为 no-op,出口 all-reduce</li><li>TP + SP: 入口处 <strong>all-gather</strong>(把 SP 切的 L 维拼回全长),出口 <strong>reduce-scatter</strong>(把全长 reduce 再切回 L 维)</li></ul>
<p>通信量:每个 transformer block 在 TP+SP 下需要 1× all-gather + 1× reduce-scatter(替换纯 TP 的 1× all-reduce),二者通信量等价,总量与纯 TP 一致。<strong>$\boxed{\text{通信量与纯 TP 完全相同}}$</strong>——但 LayerNorm / Dropout 的 activation 从 $BLD$ 降到 $BLD/T$。</p>
<div class="callout callout-info"><div class="callout-title">SP 通信量为什么和纯 TP 一样</div><p>因为 $\text{all-reduce} = \text{reduce-scatter} + \text{all-gather}$,纯 TP 末端 1× all-reduce $= $ SP 模式下 1× reduce-scatter + 后续 1× all-gathernext block 入口)。<strong>通信被重新分布,总量不变,但 activation 显存省下来了</strong></p></div>
<h3 id="83-sp-activation-memory-省了多少l3-题">8.3 SP activation memory 省了多少(L3 题)</h3>
<p><strong>未切分</strong>的单 transformer block 总 activation 内存为 $A = A_\text{in} + A_\text{out}$,其中:</p>
<ul><li>$A_\text{in}$TP <strong>内部</strong>那部分(attention 的 Q/K/V/scores、MLP 中间 [B, L, 4D]),TP 把它沿 hidden 切</li><li>$A_\text{out}$TP <strong>外部</strong>那部分(LayerNorm/Dropout/residual 的 [B, L, D] activation),纯 TP 时全副本</li></ul>
<p><strong>纯 TP 单卡 activation</strong>$A^\text{TP}_\text{per-card} = A_\text{in}/T + A_\text{out}$。</p>
<p><strong>TP + SP 单卡 activation</strong>:SP 把 TP-外那部分也沿 seq 维切(每卡持 $[B, L/T, D]$)→ $A^\text{TP+SP}_\text{per-card} = A_\text{in}/T + A_\text{out}/T = A/T$。</p>
<p><strong>SP 比纯 TP 省的部分</strong></p>
<p>$$\boxed{\;A^\text{TP}_\text{per-card} - A^\text{TP+SP}_\text{per-card} = A_\text{out} \cdot \left(1 - \frac{1}{T}\right)\;}$$</p>
<p>对 LLaMA-class 模型,$A_\text{out}$ 约占未切分总 activation 的 30-50%。TP=8 下省 $7/8 \cdot A_\text{out}$ ≈ <strong>总 activation 减少 25-40%</strong>(与 Korthikanti 2023 报告一致)。</p>
<h3 id="84-selective-activation-recompute">8.4 Selective Activation Recompute</h3>
<p>Korthikanti et al. 同篇还提了 <strong>selective recompute</strong>:只对部分 op(如 attention 的 $QK^\top$ 和 softmax 这种二次 memory 的)做 recompute,其余 op 不重算。比 full activation checkpoint 减 30-40% 计算开销,省同样多的 memory。</p>
<h2 id="9-pipeline-parallel">§9 Pipeline Parallel</h2>
<h3 id="91-朴素-pp-与-bubble">9.1 朴素 PP 与 bubble</h3>
<p>把 $L$ 层模型按深度切到 $P$ 个 stage(每 stage $L/P$ 层)。Naive PP 让每个 mini-batch 顺序流过所有 stage</p>
<pre><code>Stage 0 (GPU0): [F1] [B1]
Stage 1 (GPU1): [F1] [B1]
Stage 2 (GPU2): [F1] [B1]
Stage 3 (GPU3): [F1] [B1]
↑ 大量 GPU idle ↑ ↑大量 idle ↑</code></pre>
<p>GPU 利用率 = $1/P$,完全不能用。</p>
<h3 id="92-gpipehuang-et-al-neurips-2019">9.2 GPipeHuang et al., NeurIPS 2019</h3>
<p>把 mini-batch 切成 $M$ 个 <strong>micro-batch</strong>micro-batch 之间流水:</p>
<pre><code>Stage 0: [F1][F2][F3][F4] [B4][B3][B2][B1]
Stage 1: [F1][F2][F3][F4] [B4][B3][B2][B1]
Stage 2: [F1][F2][F3][F4] [B4][B3][B2][B1]
Stage 3: [F1][F2][F3][F4] [B4][B3][B2][B1]
|←─ &quot;warm-up&quot; ─→| |← 全部 micro-batch backward ─→| &quot;cool-down&quot;</code></pre>
<p><strong>Bubble</strong>GPU 空闲时间占比):</p>
<p>$$\boxed{\;\text{bubble ratio} = \frac{P - 1}{M + P - 1}\;}$$</p>
<p>推导:每个 micro-batch 走完 $P$ stage 要 $P$ 个 stepwarm-up 阶段(stage $i$ 等前 $i$ 个 micro-batch)共 $P-1$ 个 step idlecool-down 阶段对称 $P-1$ 个 step idle;总 step = $M + P - 1$forward+ $M + P - 1$backward= $2(M+P-1)$;其中 idle = $2(P-1)$。idle 占比 = $(P-1)/(M+P-1)$。</p>
<div class="callout callout-warn"><div class="callout-title">GPipe 的硬伤</div><p>必须等所有 $M$ 个 micro-batch 全部 forward 完才能开始 backward——所有 micro-batch 的 activation 都要存住,<strong>activation memory 与 $M$ 线性增长</strong>。这是 1F1B 解决的问题。</p></div>
<h3 id="93-1f1b--pipedreamnarayanan-sosp-2019-megatron-lm-2-sc-2021">9.3 1F1B / PipeDreamNarayanan SOSP 2019, Megatron-LM-2 SC 2021</h3>
<p><strong>1F1B = 1 Forward 1 Backward</strong>:每个 stage 在 forward 一个 micro-batch 后<strong>立刻</strong> backward(前一个已经完成的):</p>
<pre><code>Stage 0: [F1][F2][F3][F4][B1][F5][B2][F6][B3][F7][B4]...
↑ 一旦 micro-batch 1 backward 准备好(其 forward 已到 stage P
就立刻 B1,腾出 micro-batch 1 的 activation 内存</code></pre>
<p><strong>关键性质</strong></p>
<ul><li>Bubble ratio <strong>仍然是 $(P-1)/(M+P-1)$</strong>warm-up + cool-down 没省)</li><li><strong>每 stage 同时存活的 activation 数 = $P$</strong>(不是 GPipe 的 $M$)——大幅省 activation memory</li></ul>
<pre><code>Stage i 在稳态时的 activation 在内存中:
- micro-batch i forward 后等 backward 的中间 activation
- 因为有 P-i 个 stage 在 i 之后,且每 stage 走 forward / backward 都一个 step
- 所以 stage i 同时存 P-i 个 forward 完但还没 backward 的 activation
- 第一个 stage (i=0) 存最多: P 个;最后 stage 存最少: 1 个</code></pre>
<h3 id="94-interleaved-1f1bmegatron-lm-2-sc-2021">9.4 Interleaved 1F1BMegatron-LM-2, SC 2021</h3>
<p>进一步把 <strong>bubble 再压</strong>。每张 GPU 不持有连续 $L/P$ 层,而是持有 $V$ 段<strong>不连续的层</strong><strong>virtual stages</strong>):</p>
<pre><code>原 1F1B (P=4, L=8 layers):
GPU0: layers 0,1 GPU1: layers 2,3 GPU2: layers 4,5 GPU3: layers 6,7
Interleaved (P=4, V=2, L=8):
GPU0: layers 0,4 GPU1: layers 1,5 GPU2: layers 2,6 GPU3: layers 3,7
(每张 GPU 持 V=2 段, 共 L/(PV) = 1 层 / 段)</code></pre>
<p>每个 micro-batch 在 stage 0 跑 layer 0 → stage 1 跑 layer 1 → ... → stage 3 跑 layer 3 → 再回 stage 0 跑 layer 4 → ... → stage 3 跑 layer 7。<strong>一个 micro-batch 通过 stage 列 $V$ 次</strong></p>
<p><strong>Bubble ratio</strong>Narayanan et al., SC 2021, Eq. 4):</p>
<p>$$\boxed{\;\text{interleaved bubble} = \frac{P-1}{V \cdot M + P - 1} \approx \frac{P-1}{V \cdot M}\;}$$</p>
<p>把分母里 $M$ 替换成 $V \cdot M$(流水管子里同时塞了 $V$ 倍 micro-batch 的虚拟 stage 通过次数),warm-up / cool-down 的 $P-1$ 不变。代价:<strong>单 micro-batch 跨 GPU 的 send/recv 次数变 $V$ 倍</strong>(每个 micro-batch 在 stage 列上走 $V$ 圈),所以 $V$ 不能太大(一般 $V = 2$ 到 $4$)。</p>
<div class="callout callout-info"><div class="callout-title">interleaved bubble 推导直观版</div><p>普通 1F1Bbubble = $(P-1)/(M+P-1)$,分母是 micro-batch 数 $M$ 加上 warm/cool $P-1$interleaved:每个 micro-batch 跨 GPU 走 $V$ 圈,<strong>等效 micro-batch 数变成 $V \cdot M$</strong>bubble = $(P-1)/(V M + P - 1) \approx (P-1)/(VM)$。"$V$ 倍 micro-batch" 是核心直觉。</p></div>
<h3 id="95-1f1b-调度伪代码">9.5 1F1B 调度伪代码</h3>
<pre><code class="language-python">def one_f_one_b_schedule(P, M, stage_rank, num_warmup_microbatches):
&quot;&quot;&quot;
stage_rank: 当前 GPU 持有的 pipeline stage 编号 (0..P-1)
num_warmup_microbatches = P - 1 - stage_rank
&quot;&quot;&quot;
# ===== Warm-up: stage_rank 走 (P-1-stage_rank) 个 forward
for i in range(num_warmup_microbatches):
recv_activation_from_prev_stage()
out = forward(model, activation)
send_activation_to_next_stage(out)
# ===== Steady state: F1B1 alternating
num_microbatches_remaining = M - num_warmup_microbatches
for i in range(num_microbatches_remaining):
# forward
recv_activation_from_prev_stage()
out = forward(model, activation)
send_activation_to_next_stage(out)
# backward (前面 forward 过、现在到的)
recv_grad_from_next_stage()
grad_in = backward(model, grad_out)
send_grad_to_prev_stage(grad_in)
# ===== Cool-down: 剩余 backward
for i in range(num_warmup_microbatches):
recv_grad_from_next_stage()
grad_in = backward(model, grad_out)
send_grad_to_prev_stage(grad_in)</code></pre>
<h3 id="96-pp-通信特点">9.6 PP 通信特点</h3>
<p>PP 只在 stage 边界传 <strong>activation / gradient</strong>,单次传一个 micro-batch 大小 $\approx B/M \cdot L \cdot D$ bytes。<strong>通信量极小</strong>(远小于 TP / DP),所以 PP 可以跨节点(IB 也够)。</p>
<h3 id="97-pp-的硬伤load-imbalance">9.7 PP 的硬伤:load imbalance</h3>
<ul><li>不同 stage 的层数 / 计算量需手动平衡(embedding + LM head 很重,常需特殊处理)</li><li>一个 stage 的 OOM / 慢都拖死整个 pipeline(短板效应)</li><li>micro-batch 数 $M$ 太小 → bubble 大;$M$ 太大 → activation memory 涨</li></ul>
<h2 id="10-context-parallel--长序列切分">§10 Context Parallel — 长序列切分</h2>
<p>随着上下文窗口从 4K → 128K → 1M,单卡装不下完整 attention 计算。<strong>Context Parallel (CP)</strong> 沿 sequence 维切。</p>
<h3 id="101-ring-attentionliu-et-al-arxiv-231001889-2024">10.1 Ring AttentionLiu et al., arXiv 2310.01889, 2024</h3>
<p>把 $Q, K, V$ 沿 seq 维切到 $C$ 个 rank:每 rank 持 $L/C$ 个 token 的 Q/K/V。</p>
<pre><code>Rank 0: Q[0:L/C], K[0:L/C], V[0:L/C]
Rank 1: Q[L/C:2L/C], K[L/C:2L/C], V[L/C:2L/C]
...</code></pre>
<p>但 attention 需要每个 query 看到所有 key(causal 下看到所有过去 key)。<strong>Ring attention 解法</strong></p>
<ol><li>每个 rank 用本地 Q 与本地 K, V 算一块 partial attention</li><li>把本地 K, V 沿 ring 转发给下一个 rank</li><li>下一个 rank 用本地 Q × 上一 rank 的 K, V 算 partial attention,累积到 outputonline softmax 风格)</li><li>转 $C$ 步后每个 Q 都看遍全部 K, Vattention 完整</li></ol>
<p><strong>通信量</strong>:每个 rank 持 $L/C$ 个 token 的 K, V,大小 $2 \cdot L/C \cdot D$fp16 下乘 2 bytes)。Ring 转 $C-1$ 步才让每个 K, V shard 走遍所有 rank。<strong>每 rank 总传输量</strong> $\approx (C-1) \cdot 2 L D / C \approx 2 L D$(与 $C$ 几乎无关——这正是 ring 的常规性质)。</p>
<p><strong>关键</strong>:用 <strong>online softmax</strong>FlashAttention 同款)让 partial attention 可累积,不需物化完整 $L \times L$ scores。</p>
<div class="callout callout-info"><div class="callout-title">Ring attention 与 FlashAttention 关系</div><p>FlashAttention 是单卡内沿 block tiling 算 attentionblock 在 SRAM 内)。Ring attention 是把这个 tiling 推到<strong>多卡 / 多节点 level</strong>:每个 rank 持有一块 K, V,按 ring 顺序计算 partial attention 并累积。两者数学上一脉相承——都是 online softmax + block-wise accumulation。</p></div>
<h3 id="102-llama-3-cp-实现">10.2 Llama 3 CP 实现</h3>
<p>Llama 3 paper (arXiv 2407.21783) 报告在 <strong>128K 长上下文</strong> 阶段用 <strong>CP=16</strong>(短上下文 8K 阶段不需 CP,分配回 DP)。结合 FlashAttention v3128K context 单 step 时间从不可训练 → 几秒级可控。</p>
<h3 id="103-cp-与-tp-的正交性">10.3 CP 与 TP 的正交性</h3>
<ul><li>TP 沿 hidden / head 维切 → 节点内</li><li>CP 沿 sequence 维切 → 可跨节点(通信量 $\propto L$ 而非 $L^2$</li><li>复合:每张 GPU 持 $L/C$ 个 token 的 $D/T$ 维 sub-tensor</li></ul>
<h2 id="11-expert-parallel--moe-路由">§11 Expert Parallel — MoE 路由</h2>
<h3 id="111-moe-基本结构">11.1 MoE 基本结构</h3>
<p>Mixture-of-Experts 把单个 FFN 替换为 $E$ 个 expert FFN + 一个 gate / router</p>
<p>$$y = \sum_{e=1}^E G_e(x) \cdot \text{Expert}_e(x), \quad G \in \mathbb{R}^E, \quad \sum_e G_e = 1$$</p>
<p>实际部署用 <strong>top-K routing</strong>:只选 $G$ 输出最大的 $K$ 个 expert(典型 $K = 1, 2$),其余 expert 不计算。<strong>计算量与 expert 数 $E$ 无关,只与 $K$ 有关</strong>——这是 MoE 能 scale 参数的关键。</p>
<h3 id="112-expert-parallelexpert-分到不同-gpu">11.2 Expert Parallelexpert 分到不同 GPU</h3>
<table><thead><tr><th>模式</th><th>单卡 expert 数</th><th>通信</th></tr></thead><tbody><tr><td>不并行 (replicate)</td><td>$E$</td><td>0</td></tr><tr><td>TP-style expert split</td><td>每 expert 切片,全 GPU 算每个 expert</td><td>many all-reduce</td></tr><tr><td><strong>EP</strong>: 每 GPU 持 $E/N$ 个 expert</td><td>$E/N$</td><td><strong>all-to-all</strong> dispatch + combine</td></tr></tbody></table>
<p>EP 的 forward 流程:</p>
<pre><code>1. 每 GPU 对自己的 token batch 算 gate → 得到每 token 的 routing decision (assign 给 e_1, e_2, ..., e_K)
2. all-to-all dispatch: 把 token 按 expert 归属发到对应 GPU
- 输入: 每 GPU 持 B/N 个 token, 每 token 带 K 个 (expert_id, token_data)
- 输出: 每 GPU 收到所有 GPU 发给本机 expert 的 token
3. 每 GPU 用本地 expert 算 FFN
4. all-to-all combine: 把 expert 输出送回 token 原属 GPU
5. 每 GPU 按 gate weight 合并 K 个 expert 输出</code></pre>
<div class="callout callout-warn"><div class="callout-title">EP 的两个硬伤</div><p>(1) <strong>load imbalance</strong>: gate 倾向于 favor 某几个 expert,导致部分 GPU 过载;解:load balancing lossSwitch Transformer / GShard)。(2) <strong>all-to-all 通信量大</strong>: 与 token 数 × hidden 成正比,跨节点 IB 是瓶颈。DeepSeek-V3 用 <strong>node-limited routing</strong> 限制每 token 至多 dispatch 到 $M$ 个节点,减少 IB 流量。</p></div>
<h3 id="113-ep-all-to-all-代码骨架伪代码">11.3 EP all-to-all 代码骨架(伪代码)</h3>
<p>下面是 EP 前向流程的<strong>伪代码骨架</strong>——重点在 dispatch / combine 的两次 all-to-all,省去了 routing / bucketing 的工程细节(生产代码见 Megatron-Core MoE 或 DeepSpeed-MoE):</p>
<pre><code class="language-python">def expert_parallel_forward(
x, # [B, L, D]
gate, # nn.Module: tokens [BL, D] -&gt; (top_k_ids, top_k_w) 各 [BL, K]
experts_local, # 本地持有的 E_local 个 expert (nn.ModuleList)
ep_group, ep_size, ep_rank,
E_total, K,
):
&quot;&quot;&quot; 教学版:每张 GPU 持 E_local = E_total / ep_size 个 expert &quot;&quot;&quot;
B, L, D = x.shape
tokens = x.reshape(B * L, D)
# 1. routing: 每 token 选 K 个 expert
top_k_ids, top_k_w = gate(tokens) # 都是 [BL, K]
# 2. expand: top-K 把每 token 复制 K 份, 每份带一个 expert_id
expanded = tokens.unsqueeze(1).expand(-1, K, -1).reshape(B * L * K, D)
expand_ids = top_k_ids.reshape(B * L * K) # [BL*K]
expand_w = top_k_w.reshape(B * L * K)
# 3. 算每个复制 token 该送到哪个 EP rank
target_rank = expand_ids // (E_total // ep_size) # [BL*K]
# 4. 按 target_rank 排序 + 统计每 rank send_count
perm = torch.argsort(target_rank)
sorted_tokens = expanded[perm]
send_counts = torch.bincount(target_rank, minlength=ep_size).tolist()
# 5a. 交换 send_counts -&gt; recv_counts (每 rank 告诉别人自己要收多少)
send_t = torch.tensor(send_counts, dtype=torch.int64, device=x.device)
recv_t = torch.empty_like(send_t)
dist.all_to_all_single(recv_t, send_t, group=ep_group) # 一次小 a2a 同步 counts
recv_counts = recv_t.tolist()
# 5b. 同步 token 的 expert_id (用于本机分配到对应 local expert)
sorted_ids = expand_ids[perm]
received_ids = torch.empty(sum(recv_counts), dtype=sorted_ids.dtype, device=x.device)
dist.all_to_all_single(received_ids, sorted_ids,
output_split_sizes=recv_counts,
input_split_sizes=send_counts,
group=ep_group)
# 5c. 同步 token 本体
received_tokens = torch.empty(sum(recv_counts), D, device=x.device, dtype=x.dtype)
dist.all_to_all_single(received_tokens, sorted_tokens,
output_split_sizes=recv_counts,
input_split_sizes=send_counts,
group=ep_group)
# 6. local expert compute
received_out = torch.zeros_like(received_tokens)
for local_eid in range(len(experts_local)):
global_eid = ep_rank * (E_total // ep_size) + local_eid
mask = (received_ids == global_eid)
if mask.any():
received_out[mask] = experts_local[local_eid](received_tokens[mask])
# 7. all-to-all combine: 反向, 用对调过的 split sizes
combined = torch.empty_like(sorted_tokens)
dist.all_to_all_single(combined, received_out,
output_split_sizes=send_counts, # 反向
input_split_sizes=recv_counts,
group=ep_group)
# 8. 反排序 + gate weight 合并
inv_perm = torch.empty_like(perm)
inv_perm[perm] = torch.arange(perm.numel(), device=perm.device)
out_expanded = combined[inv_perm] # [BL*K, D]
out_expanded = out_expanded * expand_w.unsqueeze(-1)
out = out_expanded.view(B * L, K, D).sum(dim=1) # [BL, D]
return out.view(B, L, D)</code></pre>
<div class="callout callout-warn"><div class="callout-title">生产实现远比这复杂</div><p>真实代码处理 (a) capacity factor 限制 (避免某 expert 过载)(b) <strong>drop tokens</strong> 当 expert 满了;(c) NVL/IB hierarchical all-to-allDeepSeek 的 node-limited routing);(d) backward 的镜像 all-to-all + gate gradient。本骨架仅示意 dispatch / combine 双向流程。</p></div>
<h2 id="12-activation-memory-优化">§12 Activation Memory 优化</h2>
<h3 id="121-gradient-checkpointingchen-et-al-arxiv160406174-2016">12.1 Gradient CheckpointingChen et al., arXiv:1604.06174, 2016</h3>
<p>不存中间 activation,反向时重算。空间 $O(\sqrt{L})$(最优分段),时间 + 33% (1 次 extra forward)。</p>
<pre><code class="language-python">from torch.utils.checkpoint import checkpoint
def block_forward(x):
return transformer_block(x)
# 反向时重算 block_forward,不存中间 activation
y = checkpoint(block_forward, x, use_reentrant=False)</code></pre>
<h3 id="122-selective-recomputekorthikanti-2023">12.2 Selective RecomputeKorthikanti 2023</h3>
<p>只对<strong>二次 memory 的 op</strong>attention 的 $QK^\top$ 和 softmax)做 recompute,其余存。比 full recompute 减 30-40% 计算 overhead,省同样多 memory。Megatron-LM 默认开启。</p>
<h3 id="123-offloadzero-infinity">12.3 OffloadZeRO-Infinity</h3>
<p>把 activation / optimizer state 卸到 CPU RAM 或 NVMe。CPU 卸适合 13B-30B 量级单机训练;NVMe 卸吞吐极低,主要用于推理或 trillion-scale 大模型探索。</p>
<h3 id="124-activation-memory-公式必背">12.4 Activation Memory 公式(必背)</h3>
<p>一个 transformer block 的 activation 大致(fp16, full save):</p>
<p>$$A_\text{block} \approx \underbrace{34 \cdot B \cdot L \cdot D}_{\text{各 LayerNorm/QKV/output/MLP residual}} + \underbrace{5 \cdot B \cdot L^2 \cdot H}_{\text{attention 中间矩阵}}$$</p>
<p>详见 Korthikanti et al. 2023 Table 2。$L \gg D$ 时第二项主导(FlashAttention 直接干掉这部分);$L \approx D$ 时第一项主导(SP 把它降 $T$ 倍)。</p>
<h2 id="13-综合3d--4d--5d-parallelism">§13 综合:3D / 4D / 5D Parallelism</h2>
<p>把前面几条轴拼起来。设 world size $W$</p>
<p>$$W = D_{DP} \times T_{TP} \times P_{PP} \times C_{CP} \times E_{EP}$$</p>
<p>每条轴的角色:</p>
<table><thead><tr><th></th><th>切什么</th><th>在哪</th><th>通信原语</th></tr></thead><tbody><tr><td>DP / FSDP</td><td>batch + model states</td><td>跨节点 IB ok</td><td>all-reduce / reduce-scatter + all-gather</td></tr><tr><td>TP</td><td>单层 hidden / head</td><td><strong>节点内 NVLink</strong></td><td>all-reduce (per block × 4)</td></tr><tr><td>SP</td><td>TP 外的 LayerNorm / Dropout activation</td><td>节点内(与 TP 共 group</td><td>reduce-scatter + all-gather</td></tr><tr><td>PP</td><td>layer depth</td><td>跨节点 IB ok</td><td>point-to-point send/recv</td></tr><tr><td>CP</td><td>sequence 维</td><td>节点内或跨节点都可(通信 $\propto L$</td><td>ring K/V</td></tr><tr><td>EP</td><td>MoE expert</td><td>跨节点 IB okall-to-all 量大)</td><td>all-to-all</td></tr></tbody></table>
<h3 id="131-llama-3-405b-训练拓扑公开信息">13.1 Llama 3 405B 训练拓扑(公开信息)</h3>
<p>Meta 2024 报告 (<a href="https://arxiv.org/abs/2407.21783">2407.21783</a>)</p>
<ul><li>16K H100 GPUs(总计 16384</li><li><p>Parallelism(保持总 GPU 数不变,CP 上去时 DP 下来):</p>
<ul><li><strong>短上下文阶段 (8K)</strong>TP=8 × CP=1 × PP=16 × DP=128 = 16384</li><li><strong>长上下文阶段 (128K)</strong>TP=8 × CP=16 × PP=16 × DP=8 = 16384</li></ul></li><li>训练精度:<strong>BF16</strong>(论文 Table 报告 BF16 MFU);FP8 用于 inference 量化,<strong>未用于 405B 训练</strong></li><li><strong>54 天共 466 次中断</strong>419 unexpected + 47 planned/maintenance),其中 78% unexpected 是硬件原因</li><li>有效训练时间 &gt; 90%</li></ul>
<h3 id="132-deepseek-v3-训练拓扑公开信息">13.2 DeepSeek-V3 训练拓扑(公开信息)</h3>
<p>DeepSeek 2024 报告 (<a href="https://arxiv.org/abs/2412.19437">2412.19437</a>)</p>
<ul><li>2048 H800 GPUs</li><li>Parallelism: TP=1<strong>无 TP</strong>!靠 ZeRO + EP + PP 弥补) × PP=16 × EP=64(跨 8 节点) × ZeRO-1 DP</li><li><strong>DualPipe</strong> 双向流水(详见下节)+ all-to-all overlap</li><li>fp8 GEMM + bf16 accumulation</li></ul>
<div class="callout callout-info"><div class="callout-title">DeepSeek-V3 为什么不用 TP</div><p>V3 用了 MLA (multi-head latent attention) + 大量 MoE expertTP 切 head 收益小(latent attention head 维已经很小);而 EP 的 all-to-all overlap 配合 DualPipe 把通信全藏起来。整体 <strong>PP × EP × ZeRO</strong> 已够装下 671B 参数。</p></div>
<h2 id="14-dualpipe--2024-流水线前沿">§14 DualPipe — 2024 流水线前沿</h2>
<p>DeepSeek 2024 在 V3 paper 与独立 repo 发布的 <strong>DualPipe</strong> 算法(arXiv 2412.19437 + github.com/deepseek-ai/DualPipe)。</p>
<h3 id="141-核心想法">14.1 核心想法</h3>
<p>1F1B 的 bubble 来自 <strong>warm-up + cool-down</strong> 阶段。DualPipe 让 <strong>两个方向同时跑流水线</strong>——一组 micro-batch 从 stage 0 → P,另一组从 P → 0,两组在中间相遇时刚好填满 warm-up / cool-down 空隙。</p>
<pre><code>传统 1F1B (P=4):
Stage 0: [F1][F2][F3][F4][B1][F5][B2]...
Stage 1: [F1][F2][F3][F4][B1][F5][B2]...
(前后两端 bubble)
DualPipe (P=4):
Forward 方向 micro-batch: [F1][F2][F3][F4]...
Reverse 方向 micro-batch: ...[F4&#x27;][F3&#x27;][F2&#x27;][F1&#x27;]
Stage 0: 同时 process Forward 的 F + Reverse 的 F&#x27; + 对应 B ← 计算 / 通信完全 overlap</code></pre>
<p>更精确:DualPipe 设计了一个<strong>双向 schedule</strong>,让每张 GPU 在任意时刻都有两个 micro-batch 在 forward / backward 重叠执行;同时 expert parallel 的 all-to-all 通信也被 hide 在两组 micro-batch 的间隙里。</p>
<h3 id="142-dualpipe-vs-普通-1f1b-性质">14.2 DualPipe vs 普通 1F1B 性质</h3>
<table><thead><tr><th>维度</th><th>1F1B</th><th>Interleaved 1F1B</th><th>DualPipe</th></tr></thead><tbody><tr><td>Bubble</td><td>$(P-1)/(M+P-1)$</td><td>$(P-1)/(VM+P-1)$</td><td><strong>理想 0</strong> (warm-up / cool-down 互补)</td></tr><tr><td>Activation memory</td><td>$P$ × per stage</td><td>$V \cdot P$ ×</td><td>$\approx 2 \times P$</td></tr><tr><td>通信 overlap</td><td>部分</td><td>同 1F1B</td><td><strong>几乎 100% all-to-all overlap</strong></td></tr><tr><td>实现复杂度</td><td></td><td></td><td>极高(需 forward &amp; reverse 调度)</td></tr></tbody></table>
<h3 id="143-dualpipe-的代价">14.3 DualPipe 的代价</h3>
<ul><li><strong>代码复杂度爆炸</strong>:一个 stage 同时跑 forward (方向 1) + forward (方向 2) + backwardcuda stream 管理极困难</li><li><strong>2× activation memory</strong>:两个方向都要存 activation,比 1F1B 多 $\approx 2\times$</li><li><strong>stage 必须均衡</strong>:任何一个 stage 卡住都会让两个方向都堵</li></ul>
<p>DeepSeek 用 DualPipe 是因为 EP all-to-all 通信极重,必须想办法藏起来。<strong>一般训练任务普通 1F1B 已足够</strong></p>
<h2 id="15-torchtitan--pytorch-原生-4d-平台">§15 TorchTitan — PyTorch 原生 4D 平台</h2>
<p>Liang et al. (ICLR 2025, arXiv 2410.06511) <strong>"TorchTitan: One-stop PyTorch Native Solution for Production-Ready LLM Pre-training"</strong></p>
<h3 id="151-设计目标">15.1 设计目标</h3>
<ul><li><strong>不要再 monkey patch</strong>:把 FSDP2 / TP / PP / SP / CP / Float8 / <code>torch.compile</code> 全做进 PyTorch 主干</li><li><strong>DTensor 作为统一语言</strong>:所有切分都用 DTensor placement 描述(<code>Shard(d)</code>, <code>Replicate</code>, <code>Partial</code></li><li><strong>Composable</strong>FSDP2 与 TP 自然复合(FSDP1 与 TP 几乎不可复合)</li></ul>
<h3 id="152-代码风格与-deepspeed-monkey-patch-对比">15.2 代码风格(与 DeepSpeed monkey patch 对比)</h3>
<pre><code class="language-python"># TorchTitan 风格:声明式 DTensor placement
import torch
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
parallelize_module, RowwiseParallel, ColwiseParallel,
SequenceParallel, PrepareModuleInput,
)
from torch.distributed.fsdp import fully_shard
# 1. 构造 4D 设备 mesh
mesh = init_device_mesh(
&quot;cuda&quot;,
mesh_shape=(2, 8, 4, 8), # PP=2, FSDP=8, CP=4, TP=8
mesh_dim_names=(&quot;pp&quot;, &quot;fsdp&quot;, &quot;cp&quot;, &quot;tp&quot;),
)
# 2. 应用 TP + SP(声明每个 sub-module 的 sharding 策略)
parallelize_module(
model,
mesh[&quot;tp&quot;],
{
&quot;attn.wq&quot;: ColwiseParallel(),
&quot;attn.wk&quot;: ColwiseParallel(),
&quot;attn.wv&quot;: ColwiseParallel(),
&quot;attn.wo&quot;: RowwiseParallel(),
&quot;mlp.fc1&quot;: ColwiseParallel(),
&quot;mlp.fc2&quot;: RowwiseParallel(),
&quot;norm1&quot;: SequenceParallel(),
&quot;norm2&quot;: SequenceParallel(),
},
)
# 3. FSDP2 wrap(自动认到 mesh[&quot;fsdp&quot;]
for block in model.blocks:
fully_shard(block, mesh=mesh[&quot;fsdp&quot;])
fully_shard(model, mesh=mesh[&quot;fsdp&quot;])
# 4. PPpipeline schedule1F1B interleaved,伪代码示意)
# 真实 ScheduleInterleaved1F1B 需要 List[PipelineStage](每 rank 持 V 个 virtual stage
from torch.distributed.pipelining import PipelineStage, ScheduleInterleaved1F1B
stages = [PipelineStage(submod_v0, ...), PipelineStage(submod_v1, ...)] # V=2 virtual stages
schedule = ScheduleInterleaved1F1B(stages, n_microbatches=32, loss_fn=loss_fn)</code></pre>
<h3 id="153-float8-训练hopper--blackwell">15.3 Float8 训练(Hopper / Blackwell</h3>
<p>TorchTitan 集成 <strong>Float8 训练</strong>H100/B100 fp8 GEMM),权重 / activation 都用 fp8accumulation 在 fp32。以下是 torchao 的典型用法(具体 API 请以 torchao 当前版本文档为准):</p>
<pre><code class="language-python">import torch.nn as nn
from torchao.float8 import convert_to_float8_training, Float8LinearConfig
convert_to_float8_training(
model,
config=Float8LinearConfig(), # 默认 dynamic scaling
module_filter_fn=lambda m, n: isinstance(m, nn.Linear) and &quot;lm_head&quot; not in n,
)</code></pre>
<p>效果:H100 上 throughput +20-40%loss 几乎无差异(block-wise scaling 关键)。</p>
<h3 id="154-异步-checkpointing">15.4 异步 Checkpointing</h3>
<pre><code class="language-python">import torch.distributed.checkpoint as DCP
# 异步保存:返回 Future,不阻塞训练
future = DCP.async_save(model.state_dict(), checkpoint_id=&quot;step_10000&quot;)
# ... 继续训练
future.result() # 必要时等</code></pre>
<p>DCPDistributed Checkpoint)配合 FSDP2 的 sharded state dict<strong>单 checkpoint 写盘只需几秒</strong>(每 rank 写自己那部分到分布式存储),不阻塞训练。</p>
<h2 id="16-通信原语总结表">§16 通信原语总结表</h2>
<p>为了在面试中快速 recall,一张表收尾。</p>
<table><thead><tr><th>操作</th><th>当原始数据全部加起来是 $S$(每 rank $S$</th><th>每 rank 通信量(ring</th><th>用在</th></tr></thead><tbody><tr><td><code>all-reduce(buf, SUM)</code></td><td>$N \cdot S$ → 全 rank $S$</td><td>$2(N-1)/N \cdot S$</td><td>DDP gradient sync, TP block 末端</td></tr><tr><td><code>reduce-scatter(buf)</code></td><td>$N \cdot S$ → 各 rank $S/N$</td><td>$(N-1)/N \cdot S$</td><td>ZeRO grad reduce, SP 出口</td></tr><tr><td><code>all-gather(shard)</code></td><td>$N \cdot S/N$ → 全 rank $S$</td><td>$(N-1)/N \cdot S$</td><td>ZeRO-3 forward param, SP 入口</td></tr><tr><td><code>broadcast(buf, src)</code></td><td>rank src 的 $S$ → 全 rank $S$</td><td>$S$(树) / $(N-1)/N \cdot S$ring</td><td>model 初始化广播</td></tr><tr><td><code>all-to-all(buf)</code></td><td>$N \cdot N$ 块 → 转置</td><td>$(N-1)/N \cdot S$</td><td>MoE EP routing, sequence sharding 变换</td></tr><tr><td><code>point-to-point send/recv</code></td><td>1 → 1</td><td>$S$</td><td>PP stage 边界</td></tr></tbody></table>
<h2 id="17-25-高频面试题">§17 25 高频面试题</h2>
<p>按 L1 / L2 / L3 排序,每题 collapsible 答案要点。</p>
<h3 id="l1-必会题任何-ml-工程岗都会问">L1 必会题(任何 ML 工程岗都会问)</h3>
<details>
<summary>Q1. DDP 和 DPDataParallel)的区别?</summary>
<ul><li>DP<code>nn.DataParallel</code>):单进程多卡,主 rank scatter input + gather output<strong>有 GIL + 主 rank 瓶颈</strong>,已 deprecated</li><li>DDP<code>DistributedDataParallel</code>):多进程多卡,每进程一张 GPU,NCCL all-reduce 同步梯度</li><li>DDP 比 DP 快 1.5-3 倍且 scaling 远好</li></ul>
<p>陷阱:说 DP 是"训练快"——错,DP 是历史遗留 API,生产用 DDP。</p>
</details>
<details>
<summary>Q2. NCCL all-reduce 的通信量?</summary>
<ul><li>Ring 算法:单 GPU 总流量 $2(N-1)/N \cdot S \approx 2S$ bytes</li><li>与 $N$ 几乎无关(这是 ring 的精髓)</li><li>等价于 reduce-scatter ($S$) + all-gather ($S$)</li></ul>
<p>陷阱:说 $N \cdot S$ 或 $S/N$;忘了 ring 的 step 数和 per-step 流量。</p>
</details>
<details>
<summary>Q3. Adam mixed-precision 训练单卡 model states 多少?</summary>
<ul><li>参数 fp16: $2\Phi$</li><li>梯度 fp16: $2\Phi$</li><li>Optimizer (fp32 master + Adam m + v): $4\Phi + 4\Phi + 4\Phi = 12\Phi$</li><li><strong>总计 $16\Phi$ bytes</strong></li></ul>
<p>陷阱:忘了 fp32 master copy;或把 Adam $m, v$ 当 fp16(实际 fp32)。</p>
</details>
<details>
<summary>Q4. ZeRO 1/2/3 分别切什么?</summary>
<ul><li>ZeRO-1: 切 optimizer state</li><li>ZeRO-2: + gradient</li><li>ZeRO-3: + parameter(最激进)</li></ul>
<p>单卡 model states 从 $16\Phi$ 降到 $16\Phi/N$ZeRO-3)。</p>
<p>陷阱:把顺序搞反;不知道 ZeRO-1/2 通信量与 DDP 完全一样(都是 $2\Phi$)。</p>
</details>
<details>
<summary>Q5. FSDP 和 ZeRO-3 有什么区别?</summary>
<ul><li>算法层面:<strong>完全等价</strong>FSDP 就是 ZeRO-3 的 PyTorch 实现)</li><li>工程层面:FSDP 集成在 PyTorch 主干;ZeRO-3 在 DeepSpeed 库</li><li>FSDP2 改用 per-parameter DTensor sharding,与 TP 复合自然,state_dict 简洁</li></ul>
<p>陷阱:说"FSDP 比 ZeRO 通信少"或反之——错,通信量同。</p>
</details>
<details>
<summary>Q6. Tensor Parallel 切 attention 的哪一维?</summary>
<ul><li><strong>head 维</strong>(每张卡持 $H/T$ 个 head</li><li>$W_Q, W_K, W_V$ column-parallel(按输出维切)</li><li>$W_O$ row-parallel(按输入维切)</li><li>每个 attention block forward 1× all-reduce, backward 1× all-reduce</li></ul>
<p>陷阱:说切 hidden_dim 维;或忘了 col + row 配对让中间不通信。</p>
</details>
<details>
<summary>Q7. TP 必须放在哪里?</summary>
<ul><li><strong>节点内 NVLink 域</strong>(带宽 900 GB/s</li><li>跨节点 IB(50 GB/s)走 TP 性能会暴跌</li><li>这就是为什么 TP=8 是黄金尺寸(一个节点 8 卡)</li></ul>
<p>陷阱:说 TP 可以任意跨节点;或把节点内 / 跨节点带宽搞反。</p>
</details>
<details>
<summary>Q8. Pipeline Parallel 的 bubble 怎么算?</summary>
<ul><li>朴素 PP$M = 1$bubble = $(P-1)/P$(巨大)</li><li>GPipe / 1F1B with $M$ micro-batch$(P-1)/(M+P-1)$</li><li>通常 $M \geq 4P$ 让 bubble &lt; 20%</li></ul>
<p>陷阱:忘了 $M$ 在分母;只说 $1/P$ 不说 $M$。</p>
</details>
<details>
<summary>Q9. 1F1B 比 GPipe 好在哪?</summary>
<ul><li><strong>bubble ratio 一样</strong> $(P-1)/(M+P-1)$</li><li>但 1F1B 每 stage 同时存活的 activation 数 = $P$(不是 GPipe 的 $M$</li><li>大幅省 <strong>activation memory</strong></li></ul>
<p>陷阱:说 1F1B 减小了 bubble——错,1F1B 不减 bubble,省 activation。</p>
</details>
<details>
<summary>Q10. 激活检查点(gradient checkpointing)的代价?</summary>
<ul><li>显存:$O(L) \to O(\sqrt{L})$</li><li>时间:+ 33%(一次额外 forward</li><li>实际生产几乎必开(70B+ 模型不开就 OOM)</li></ul>
<p>陷阱:说"显存减半"——不准确,理论是 $\sqrt{L}$;或说"时间减半"。</p>
</details>
<h3 id="l2-进阶题research-oriented-岗位">L2 进阶题(research-oriented 岗位)</h3>
<details>
<summary>Q11. ZeRO-3 forward 为什么需要 all-gather</summary>
<ul><li>参数被切到 $N$ 张卡 → 每张卡只持 $1/N$</li><li>算某层 forward 时需要完整 $W^{(\ell)}$ → 临时 <strong>all-gather</strong> 到所有 rank</li><li>forward 完立刻 <strong>release</strong>,只保留 shard</li><li>backward 时再次 all-gatherforward 中已释放)</li></ul>
<p>陷阱:以为参数在所有 rank 上常驻;或忘了 backward 也要重新 all-gather。</p>
</details>
<details>
<summary>Q12. 推导:DDP vs ZeRO-3 通信量差。</summary>
<ul><li>DDP: backward 1× all-reduce = $2\Phi$,总 $2\Phi$</li><li>ZeRO-3: forward 1× all-gather ($\Phi$) + backward 1× all-gather ($\Phi$) + 1× reduce-scatter ($\Phi$) = $3\Phi$</li><li><strong>ZeRO-3 比 DDP 多 50% 通信,换 $N\times$ 显存下降</strong></li></ul>
<p>陷阱:以为 ZeRO 减通信——错,ZeRO 减显存,可能增通信(视阶段)。</p>
</details>
<details>
<summary>Q13. NCCL ring all-reduce 单步流量?</summary>
<ul><li>总 step 数:$2(N-1)$reduce-scatter $N-1$ + all-gather $N-1$</li><li>单 step:每 rank 发送 $S/N$ bytes</li><li>单 GPU 总流量:$2(N-1) \cdot S/N \approx 2S$</li><li>Bandwidth 利用率 $\to 1$ 当 $N \to \infty$</li></ul>
<p>陷阱:说 $S$ 而非 $2S$;或说 $N \cdot S$(忘了 ring 的精髓)。</p>
</details>
<details>
<summary>Q14. SPSequence Parallel)省了什么?</summary>
<ul><li>不省通信(与纯 TP 通信量一样)</li><li><strong>TP 外的 activation memory</strong>LayerNorm / Dropout</li><li>把全副本 $[B, L, D]$ 切成 $[B, L/T, D]$</li><li>总 activation memory 降 25-40%</li></ul>
<p>陷阱:说"SP 减通信"——错,SP 只是把 all-reduce 重排成 reduce-scatter + all-gather(总量等同),但 activation 切了。</p>
</details>
<details>
<summary>Q15. Interleaved 1F1B 怎么减 bubble</summary>
<ul><li>每 stage 持 $V$ 个 virtual stage(不连续的 $V$ 段 layer</li><li>bubble: $(P-1)/(VM+P-1) \approx (P-1)/(VM)$</li><li>同 $M$ 下 bubble 降 $V$ 倍</li><li>代价:通信次数 × $V$</li></ul>
<p>陷阱:说"interleaved 减计算量"——错,只重排时间轴;或忘了通信 × V 的代价。</p>
</details>
<details>
<summary>Q16. MoE 的 EP all-to-all 通信量?</summary>
<ul><li>每 token 选 $K$ 个 expert</li><li>Dispatch: 每 GPU 把本机 $B/N$ token 发到对应 expert 的 GPU</li><li>Combine: 反向</li><li>总 per-rank 通信 $\approx 2 \cdot K \cdot B/N \cdot D$(双向 all-to-all</li></ul>
<p>陷阱:忘了 top-K(不是 $E$);忘了 dispatch + combine 两次。</p>
</details>
<details>
<summary>Q17. HSDP 和 FSDP 区别?</summary>
<ul><li>FSDP:所有 GPU 同一 sharding group</li><li>HSDP:组内 FSDP / ZeRO-3,组间 DDP</li><li>HSDP 单 group 内通信少(节点内 NVLink),组间用 grad all-reduce(不需 weight all-gather</li><li>trade-off:组内 model states 多(不切到全 world),换跨节点通信减</li></ul>
<p>陷阱:以为 HSDP 减总通信——精度上是减跨节点通信,组内 / 组间是 trade-off。</p>
</details>
<details>
<summary>Q18. Llama 3 405B 用了哪些并行?</summary>
<ul><li>16K H100(总 16384,保持不变;CP 上去时 DP 下来)</li><li>短 context (8K)TP=8 × CP=1 × PP=16 × DP=128</li><li>长 context (128K)TP=8 × CP=16 × PP=16 × DP=8</li><li>54 天共 466 次中断(419 unexpected + 47 planned),90%+ 有效训练时间</li><li>训练用 BF16(不是 FP8——FP8 是推理量化)</li><li>Meta 2024, arXiv 2407.21783</li></ul>
<p>陷阱:说用了 EP——错,Llama 3 是 dense 模型不用 EP;或忘了 CP 这一阶段的存在。</p>
</details>
<details>
<summary>Q19. ZeRO++ 三个 trick</summary>
<ul><li><strong>qwZ</strong>forward all-gather 用 int8 量化(block-wise quant</li><li><strong>hpZ</strong>weight 在节点内复制(NVLink 域),跨节点仍切——backward all-gather 走 NVLink</li><li><strong>qgZ</strong>backward gradient reduce-scatter 也走 int8</li><li>总通信量降 4×,throughput 384 GPU 上 +116% (Wang 2023)</li></ul>
<p>陷阱:以为 ZeRO++ 改训练精度——只在通信途中量化 buffer,权重和计算仍是 fp16/bf16。</p>
</details>
<details>
<summary>Q20. 一个 7B 模型在 8 张 A100-40G 上能跑 fp16 训练吗?</summary>
<ul><li>Model states (Adam): $16 \times 7$B $= 112$ GB</li><li>DDP: 单卡 112 GB<strong>单卡放不下</strong>A100 40G</li><li>ZeRO-3 / FSDP: 单卡 $112/8 = 14$ GB ✓</li><li>加 activation (with checkpoint): 几 GB ✓</li><li>加 workspace: 几 GB</li><li><strong>结论:FSDP/ZeRO-3 + 激活检查点可以跑</strong></li></ul>
<p>陷阱:说"DDP 也能跑";或忘了 model states 而只算参数 $\Phi$;或忘了 fp32 master copy。</p>
</details>
<h3 id="l3-高级变体顶级-lab--自研基建岗">L3 高级变体(顶级 lab / 自研基建岗)</h3>
<details>
<summary>Q21. 推导 1F1B bubble ratio + interleaved 怎么进一步减小。</summary>
<ul><li>1F1Bwarm-up $P-1$ step (forward fill) + cool-down $P-1$ step (backward drain) + steady $M-P+1$ step</li><li>总 step $= 2M$forward + backward)需要 $2(M + P - 1)$ 时间槽(含 warm/cool 各 $P-1$ idle 槽)</li><li>bubble = $2(P-1) / [2(M+P-1)] = (P-1)/(M+P-1)$</li><li><strong>Interleaved 1F1B</strong>Narayanan SC 2021, Eq. 4):每 GPU 持 $V$ 段不连续 layer (virtual stages),每 micro-batch 在 stage 列上走 $V$ 圈 → 流水里等效 $V \cdot M$ 个 micro-batch 在排队</li><li>bubble = $(P-1) / (V \cdot M + P - 1) \approx (P-1)/(V \cdot M)$</li><li><strong>同 $M$ 下 bubble 降 $V$ 倍;代价:跨 GPU send/recv 次数 $\times V$</strong></li></ul>
<p>陷阱:忘了 interleaved 通信 × V 的代价;说"interleaved 减 forward 计算"——只重排时间。</p>
</details>
<details>
<summary>Q22. FSDP vs ZeRO-3 通信量 + 各自适用场景。</summary>
<ul><li><strong>通信量完全一样</strong>forward all-gather $\Phi$ + backward all-gather $\Phi$ + reduce-scatter $\Phi$ = $3\Phi$</li><li><p>工程差异:</p>
<ul><li>FSDP2: 用 DTensor 描述 per-param sharding,与 TP / PP 复合自然,<code>torch.compile</code> 友好</li><li>DeepSpeed ZeRO-3: 用 flat-parameter,要 monkey patch 拼装,但 ZeRO++ / Offload / Infinity 生态完整</li></ul></li><li><p>选型:</p>
<ul><li>新项目(dense + 4D 并行)→ <strong>FSDP2 / TorchTitan</strong></li><li>MoE + 跨框架 + offload → <strong>DeepSpeed</strong></li></ul></li></ul>
<p>陷阱:说"FSDP 比 ZeRO 通信少"——算法相同;或说"FSDP 不支持 offload"——FSDP2 OffloadPolicy 已支持。</p>
</details>
<details>
<summary>Q23. TP + SP 相比纯 TP 省多少 activation memory</summary>
<ul><li>设 transformer block activation $A_\text{block}$</li><li>TP-内 (attention 中间 / MLP 中间):约 $A_\text{block} \times 0.5-0.7$TP 切了</li><li>TP-外 (LayerNorm / Dropout / residual):约 $A_\text{block} \times 0.3-0.5$<strong>TP 不切(全副本)</strong></li><li>纯 TP 单卡 activation = $A_\text{TP-内}/T + A_\text{TP-外}$</li><li>TP+SP 单卡 activation = $A_\text{TP-内}/T + A_\text{TP-外}/T = A_\text{block}/T$</li><li><strong>省了 $A_\text{TP-外} \cdot (1 - 1/T)$</strong>,约总 activation 的 <strong>25-40%</strong>(取决于模型)</li></ul>
<p>陷阱:说"SP 让 activation 减 $T$ 倍"——不准确,只对 TP-外那部分;或忘了 SP 通信量与纯 TP 一样。</p>
</details>
<details>
<summary>Q24. DualPipe 相对 1F1B 的核心改进?</summary>
<ul><li><strong>双向流水线</strong>:一组 micro-batch 从 stage 0 → P,另一组从 P → 0</li><li>两组在中间相遇时刚好填满 1F1B 的 warm-up / cool-down bubble</li><li><strong>理论 bubble = 0</strong>(理想下两侧互补)</li><li>关键收益:<strong>all-to-all 通信完全 overlap</strong>EP 重要)</li><li>代价:activation memory $\times 2$(两个方向都要存);实现极复杂</li><li>DeepSeek-V3 December 2024 报告 (arXiv 2412.19437) + github.com/deepseek-ai/DualPipe</li></ul>
<p>陷阱:说"DualPipe 是新的 PP 切法"——它是新的 schedule;或忘了 2× activation 的代价。</p>
</details>
<details>
<summary>Q25. 现在 frontier 训练(如 DeepSeek-V3 / Llama 3)的 4D / 5D parallelism 怎么组合?</summary>
<ul><li><strong>5 个正交维度</strong>DP / FSDP × TP × PP × CP × EP</li><li>World size $W = D \times T \times P \times C \times E$</li><li><p>经验法则:</p>
<ul><li>TP=8(一个节点 NVLink 域内必须塞下)</li><li>PP=8-32(跨节点 IB okbubble 控制 $M \geq 4P$</li><li>CP 看 context 长度(Llama 3 在 128K 用 CP=161M 可能 CP=64+</li><li>EP 看 MoE expert 数(DeepSeek-V3 用 EP=64</li><li>FSDP / DP 用剩下的 world size</li></ul></li><li><p><strong>Llama 3 405B</strong>(保持总 16K GPU 不变,CP 上 DP 下):</p>
<ul><li>8K context: TP=8 × CP=1 × PP=16 × DP=128</li><li>128K context: TP=8 × CP=16 × PP=16 × DP=8</li></ul></li><li><strong>DeepSeek-V3 671B</strong>: TP=1 + PP=16 × EP=64 × ZeRO-1 DP=2,共 2048 H800</li><li><p><strong>关键工程点</strong></p>
<ul><li>通信原语按拓扑放(TP 走 NVLinkPP / DP 走 IB</li><li>DualPipe / Interleaved 1F1B 压 PP bubble</li><li>FlashAttention v3 + SP / CP 压 activation</li><li>fp8 GEMMH100 / B100+ bf16 accumulation 提 throughput</li></ul></li></ul>
<p>陷阱:把维度顺序搞反;忘了 TP 必须节点内;以为 DeepSeek 用了 TPV3 实际 TP=1)。</p>
</details>
<h2 id="a-附录完整-4d-wrap-代码骨架">§A 附录:完整 4D wrap 代码骨架</h2>
<p>下面是一个 4D parallelism 的 minimal 端到端 wrap 示例(FSDP2 + TP + SP + PP),按 TorchTitan 风格。</p>
<pre><code class="language-python">import torch
import torch.nn as nn
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.parallel import (
parallelize_module, RowwiseParallel, ColwiseParallel,
SequenceParallel, PrepareModuleInput, PrepareModuleOutput,
)
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
from torch.distributed.pipelining import pipeline, SplitPoint, ScheduleInterleaved1F1B
from torch.utils.checkpoint import checkpoint
def build_4d_model(model, world_size):
&quot;&quot;&quot; 4D parallelism wrap (PP × FSDP × CP × TP) &quot;&quot;&quot;
# Step 1. 建 4D device mesh
# 例: 64 GPU = PP=2 × FSDP=4 × CP=1 × TP=8
mesh = init_device_mesh(
&quot;cuda&quot;,
mesh_shape=(2, 4, 1, 8),
mesh_dim_names=(&quot;pp&quot;, &quot;fsdp&quot;, &quot;cp&quot;, &quot;tp&quot;),
)
# Step 2. TP + SP wrap (mesh[&quot;tp&quot;])
tp_plan = {
# Attention
&quot;attn.wq&quot;: ColwiseParallel(),
&quot;attn.wk&quot;: ColwiseParallel(),
&quot;attn.wv&quot;: ColwiseParallel(),
&quot;attn.wo&quot;: RowwiseParallel(),
# MLP
&quot;mlp.fc1&quot;: ColwiseParallel(),
&quot;mlp.fc2&quot;: RowwiseParallel(),
# SP: norm / residual 沿 seq 维切
&quot;ln1&quot;: SequenceParallel(),
&quot;ln2&quot;: SequenceParallel(),
}
for block in model.blocks:
parallelize_module(block, mesh[&quot;tp&quot;], tp_plan)
# Step 3. Activation checkpoint (每 block 一次)
for i, block in enumerate(model.blocks):
model.blocks[i] = _ckpt_wrap(block)
# Step 4. PP split (mesh[&quot;pp&quot;])
# 把 model 切成 2 段(PP=2
split_spec = {&quot;blocks.16&quot;: SplitPoint.BEGINNING}
pipe = pipeline(
model,
mb_args=(torch.randn(1, 4096, device=&quot;cuda&quot;),),
split_spec=split_spec,
)
pp_stage = pipe.build_stage(
stage_index=mesh[&quot;pp&quot;].get_local_rank(),
device=torch.device(f&quot;cuda:{torch.cuda.current_device()}&quot;),
group=mesh[&quot;pp&quot;].get_group(),
)
# Step 5. FSDP2 wrap (mesh[&quot;fsdp&quot;])
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
)
for module in pp_stage.submod.modules():
if isinstance(module, TransformerBlock):
fully_shard(module, mesh=mesh[&quot;fsdp&quot;], mp_policy=mp_policy)
fully_shard(pp_stage.submod, mesh=mesh[&quot;fsdp&quot;], mp_policy=mp_policy)
return pp_stage, mesh
def _ckpt_wrap(block):
&quot;&quot;&quot; 把一个 TransformerBlock 包成 activation checkpoint 形式 &quot;&quot;&quot;
class Wrapped(nn.Module):
def __init__(self, inner):
super().__init__()
self.inner = inner
def forward(self, *args, **kw):
return checkpoint(self.inner, *args, use_reentrant=False, **kw)
return Wrapped(block)
# ============================================================
# Training loop with PP schedule
# ============================================================
def train_step_4d(pp_stage, schedule, optimizer, batch):
&quot;&quot;&quot; 1 step of 4D parallel training &quot;&quot;&quot;
optimizer.zero_grad(set_to_none=True)
# PP schedule runs forward + backward across all stages
losses = []
schedule.step(batch, losses=losses) # 内部触发 FSDP all-gather / TP all-reduce / etc.
optimizer.step()
return torch.stack(losses).mean()</code></pre>
<p><strong>说明</strong>:上述代码是教学骨架,实际生产请直接用 TorchTitan repopytorch/torchtitan)——其包含完整的 <code>Trainer</code>、checkpointing、profiling、loss / lr schedule,本节只给概念示意。</p>
<h2 id="b-参考资料">§B 参考资料</h2>
<ul><li><strong>DDP</strong>: PyTorch documentation, <code>nn.parallel.DistributedDataParallel</code></li><li><strong>ZeRO</strong>: Rajbhandari et al., "ZeRO: Memory Optimizations Toward Training Trillion Parameter Models", SC 2020 (<a href="https://arxiv.org/abs/1910.02054">arXiv:1910.02054</a>)</li><li><strong>ZeRO-Offload</strong>: Ren et al., USENIX ATC 2021</li><li><strong>ZeRO-Infinity</strong>: Rajbhandari et al., SC 2021</li><li><strong>ZeRO++</strong>: Wang et al., "ZeRO++: Extremely Efficient Collective Communication for Giant Model Training", arXiv:2306.10209, 2023</li><li><strong>FSDP</strong>: Zhao et al., "PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel", VLDB 2023</li><li><strong>FSDP2 / TorchTitan</strong>: Liang et al., "TorchTitan: One-stop PyTorch Native Solution", ICLR 2025 (<a href="https://arxiv.org/abs/2410.06511">arXiv:2410.06511</a>)</li><li><strong>Megatron-LM TP</strong>: Shoeybi et al., "Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism", arXiv:1909.08053, 2019</li><li><strong>Megatron-LM-2 (Interleaved 1F1B)</strong>: Narayanan et al., "Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM", SC 2021</li><li><strong>GPipe</strong>: Huang et al., "GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism", NeurIPS 2019</li><li><strong>PipeDream / 1F1B</strong>: Narayanan et al., "PipeDream: Generalized Pipeline Parallelism for DNN Training", SOSP 2019</li><li><strong>Sequence Parallel + Selective Recompute</strong>: Korthikanti et al., "Reducing Activation Recomputation in Large Transformer Models", MLSys 2023</li><li><strong>Ring Attention</strong>: Liu et al., "Ring Attention with Blockwise Transformers for Near-Infinite Context", ICLR 2024 (<a href="https://arxiv.org/abs/2310.01889">arXiv:2310.01889</a>)</li><li><strong>Gradient Checkpointing</strong>: Chen et al., "Training Deep Nets with Sublinear Memory Cost", arXiv:1604.06174, 2016</li><li><strong>DualPipe</strong>: DeepSeek-V3 Technical Report, <a href="https://arxiv.org/abs/2412.19437">arXiv:2412.19437</a>, December 2024</li><li><strong>Llama 3</strong>: Meta AI, "The Llama 3 Herd of Models", <a href="https://arxiv.org/abs/2407.21783">arXiv:2407.21783</a>, 2024</li><li><strong>GShard / Switch Transformer (MoE EP)</strong>: Lepikhin et al. 2020 / Fedus et al. JMLR 2022</li></ul>
<footer class="aris-footer">
Generated by <a href="https://github.com/wanshuiyin/Auto-claude-code-research-in-sleep/blob/main/skills/render-html/SKILL.md">ARIS <code>/render-html</code></a> ·
source path <code>docs/tutorials/distributed_training_tutorial.md</code> ·
SHA256 <code>86087779b70e</code> ·
generated at 2026-05-19 05:38 UTC.
This is a generated view — edit the source Markdown, then re-render.
</footer>
</main>
</div>
</body>
</html>