关卡阅读材料 // 06
Transformer 扩展
把低精度数值、层内切分和通信重叠组合到一个 Transformer 训练步骤中,理解数值范围与并行维度必须一起设计。
建议边读边写下答案;时长包含代码推演与练习。
进入互动关卡→- 01导读与目标05 MIN
- 02心智模型08 MIN
- 03概念深挖15 MIN
- 04代码推演12 MIN
- 05历史生态08 MIN
- 06练习复盘12 MIN
把 CUDA Python、行业框架与 Profiling 工具组合进可复现工作流,而不是只优化孤立 Kernel。
关卡把 FP8 Recipe、Tensor Parallel 与通信重叠放在同一训练步骤;实验要求同时验证数值误差、通信依赖和时间线。
- 课程 03–04
- PyTorch 基础
- 理解 Transformer Linear 与反向传播
- PyTorch
- Transformer Engine
- NCCL
- Nsight Systems / Compute
导读与目标
先记住这一句
先让数值落入可表示范围,再切分正确的维度,只重叠真正独立的工作。
完成后你应该能够
- 01
解释 FP8 格式、Scale、amax 与高精度主状态的职责。
- 02
为宽层、深模型、长序列和大 Batch 分别选择合适的并行轴。
- 03
画出 Tensor Parallel 中计算与 Collective 的依赖边。
- 04
判断一次异步通信是否拥有真实且安全的重叠窗口。
建立心智模型
- 01缩放
用 amax 等统计量计算 Scale,把张量映射进 FP8 的有限表示范围。
- 02切层
Tensor Parallel 切分单层权重;Data、Pipeline 与 Context Parallel 分别处理 Batch、深度与序列。
- 03重叠
Collective 发出后,只有不依赖其结果的计算才能安全执行;消费者之前必须等待。
概念深挖
E4M3 与 E5M2 有分工
E4M3 保留更多尾数精度;E5M2 用更少尾数换取更大指数范围。HYBRID Recipe 常让前向与反向使用不同格式。
并行轴解决不同问题
TP 面向宽层,PP 面向深度,CP 面向长序列激活,DP 面向 Batch;选择错误的轴不会解决原瓶颈。
更多 TP 也会增加通信
局部 GEMM 变小后,可隐藏通信的计算窗口可能同时缩短,因此并行度不是越大越好。
低精度首先是范围管理
E4M3 与 E5M2 只有 8 位,能够表示的精度与范围都有限。Scale 把原始张量映射到格式可表示区间,amax 是估计范围的常用统计量。若 Scale 太小,大值饱和;太大则许多小值被压到相同量化值或零。
训练分布会随步骤变化,Delayed 或 Current Scaling Recipe 以不同方式更新 Scale。HYBRID 常让前向权重与激活用 E4M3、反向梯度用范围更大的 E5M2。并非所有算子都适合 FP8,高精度主权重和部分归约仍可能必要。
停下来想一想为什么固定一次 Scale 后训练到结束风险很高?+
参考答案张量分布会漂移;固定 Scale 可能逐渐导致饱和或低值分辨率不足。需要跟踪统计并验证稳定性。
并行轴是一组正交问题
Data Parallel 切 Batch 并复制模型;Tensor Parallel 切单层张量;Pipeline Parallel 按层深度分 Stage;Context Parallel 切序列与相关激活。它们可以组合,但每增加一个轴都会引入新的通信、调度或状态管理。
选择应从瓶颈出发:宽层 OOM 看 TP,深模型分层看 PP,长序列激活看 CP,吞吐扩展看 DP。总 GPU 数常是多个并行度的乘积,Rank 分组必须让每种 Collective 只发生在正确的通信域。
停下来想一想长上下文导致 Activation OOM,只增加 DP 有帮助吗?+
参考答案通常没有直接帮助,因为每个 DP Rank 仍处理自己的完整序列;可研究 Context Parallel、Activation Checkpoint 或序列切分。
异步只是句柄,独立工作才是窗口
把 All-Gather 改成 async 会立即返回 Handle,但如果下一行马上 wait,通信仍完全暴露。要重叠,程序必须找到不读取 Gather 结果、又能在同一时间执行的计算,并确保它不会争用到让两者互相拖慢。
Tensor Parallel 度增加时,每个 Rank 的 GEMM 变小,而 Collective 参与者和通信占比可能上升,重叠窗口反而缩短。应在 Profiler 中确认计算与通信时间线,检查等待点、Stream、消息大小和拓扑。
停下来想一想all_gather_async 后立即 handle.wait() 与同步 All-Gather 有何性能差异?+
参考答案通常很小,因为没有独立工作填入两者之间;异步接口本身没有创造重叠。
把概念放回代码
01recipe = DelayedScaling(fp8_format=Format.HYBRID)02with te.autocast(enabled=True, recipe=recipe):03 local = column_parallel_linear(x)04 handle = all_gather_async(local)05 independent = other_linear(x)06 gathered = handle.wait()代码逐段推演
- 01RECIPE声明格式与 Scale 策略
Recipe 把格式选择、统计历史和 Scale 更新集中管理,仍需模型级精度与收敛验证。
- 02LOCAL每个 TP Rank 计算局部 Linear
列切或行切决定哪一维属于当前 Rank,也决定后续需要 Gather、Reduce 或 Reduce-Scatter。
- 03ASYNC发起通信并保留 Handle
此时结果尚不可安全消费;Buffer 生命周期与通信 Stream 必须保持有效。
- 04WAIT只在消费者边界等待
把真正独立的 other_linear 放入窗口,随后在首次读取 Gather 结果前 wait。
知识图谱
补充参考 · 不要求一次读完,也不计入核心 60 分钟
核心术语
- FP8
- 8 位浮点格式家族;较小表示范围要求明确的缩放策略和高精度累加/保留路径。
- amax
- 张量或历史窗口中的最大绝对值,用于估计缩放因子和检测溢出风险。
- Delayed Scaling
- 使用历史 amax 估计当前 Scale,以减少量化前额外读取;Scale 滞后是其精度交换。
- Tensor Parallel
- 沿模型张量维度切分 Linear/Attention 等计算,每个 Rank 计算局部块并通过 Collective 组合。
- Overlap
- 让通信与真正不依赖其结果的计算同时推进;异步调用本身并不等于有效重叠。
带数字推演
- 01计算 Scale
scale ≈ 448 ÷ 240 ≈ 1.87;原值 120 会映射到约 224。
- 02观察突增
若新值突然到 300,会映射到约 560,超过 448 后发生裁剪。
- 03计算步骤
T_step ≈ 12 + 5 − 3 = 14 ms;未被覆盖的通信仍在关键路径上。
诊断手册
FP8 Loss 突增或出现 NaN
- 先检查
- amax/Scale 历史、裁剪、敏感层、主权重与归约精度
- 需要的证据
- 逐层 amax、溢出计数、BF16 对照和多步 Loss 曲线
TP 增大后步骤更慢
- 先检查
- 局部 GEMM 是否过小、Collective 占比、拓扑和同步边界
- 需要的证据
- 按 Rank 的计算/通信区间与 TP 度扫描
异步 Collective 没有隐藏
- 先检查
- 后续工作是否真的独立、是否共用执行资源、wait 是否过早
- 需要的证据
- Nsight Systems 中 Collective、GEMM 与 wait 的重叠区间
从硬件到生态
- 01FP32 训练
统一高精度简化数值推理。
显存、带宽与矩阵吞吐成本较高。 - 02FP16/BF16 混合精度
低精度 Tensor 路径配合高精度累加与 Loss Scaling。
数值策略成为性能工程的一部分。 - 03Hopper FP8 + Transformer Engine
硬件格式、Scale Recipe 与框架模块协同。
优化对象从单算子扩展为跨层统计、并行组与通信调度。
硬件与生态坐标
Hopper Tensor Core 加入 FP8 矩阵路径;Transformer Engine 在软件层管理 FP8-safe 算子、Scale 与 amax 历史。Megatron Core 等框架继续把 TP、PP、DP、CP 与通信重叠组合成可配置的训练系统。
练习与复盘
超宽的单个 Linear 层放不进一张 GPU,应优先研究哪个并行轴?
- A. Tensor Parallel
- B. 只增加 Data Parallel
- C. 只把 Batch 改成 1
查看答案+
Tensor Parallel 直接切分层内权重与计算;Data Parallel 仍会在每个 Rank 复制完整模型层。
为什么反向梯度常更偏向 E5M2,而前向常用 E4M3?
提示+
比较指数范围与尾数精度。
参考答案+
梯度可能需要更大动态范围,E5M2 用更少尾数换更大指数范围;前向权重与激活更重视精度,常使用 E4M3。
模型很深、单层能放下,但整体模型放不进一张卡。优先研究哪个轴?
提示+
按层深度切分。
参考答案+
Pipeline Parallel,把不同层分配到不同 Stage;同时还要处理 Pipeline Bubble、Microbatch 与跨 Stage 激活通信。
异步 All-Gather 发出后,另一个 Linear 读取 Gather 输出的一部分。能否重叠?
提示+
它是否真的独立?
参考答案+
不能直接重叠,因为存在读后依赖。除非算法支持分块消费并建立更细粒度完成信号,否则必须等待对应数据完成。
TP 从 4 增到 8 后变慢,应优先比较什么?
提示+
局部计算变小,通信域变大。
参考答案+
比较每 Rank GEMM Shape/效率、Collective 消息与耗时、重叠窗口、Rank 拓扑和等待比例,再判断是否 TP 过度。
可选实践实验
不计入核心 60 分钟 · 需要相应 CUDA / GPU 环境
实验目标
验证一次 FP8 + Tensor Parallel 训练步骤
在支持环境中比较 BF16 与 FP8 Recipe,并用时间线检查 TP Collective 是否与真正独立的计算重叠。建议步骤
- 01
建立 BF16 参考步骤,固定随机种子、输入 Shape、损失与梯度摘要。
- 02
启用 Transformer Engine Recipe,记录格式、Scale 策略、支持硬件和不能进入低精度的算子。
- 03
运行 TP=1 与 TP>1,保存每 Rank Shape、Collective 类型和结果误差。
- 04
用 Profiler 检查异步 Collective、独立 Linear 与 wait 边界;若无重叠,解释依赖或资源原因。
提交环境能力、BF16/FP8 数值对比、TP Rank 映射与时间线。没有支持 FP8 的硬件时,完成 Recipe/Shape 静态推演并明确限制。
改变 TP 度,画出局部 GEMM 时间、Collective 时间与重叠比例,找出过度切分的拐点。