跳转至

七:训练加速技巧

什么是算子融合 (Operator Fusion)?在 Transformer 中,常见的融合模式有哪些?

算子融合是深度学习框架或编译器将多个连续的计算操作合并成一个单一的、更高效的计算内核。它的核心动机是减少显存带宽的消耗。在GPU上,计算速度远超数据搬运速度(“内存墙”)。如果每个小操作都独立执行,每次都需要将数据从显存(HBM)读取到计算单元,完成后再写回显存,下一个操作再重复这个过程。大量时间浪费在数据搬运上,而不是实际计算。

融合后的算子只需从显存读取一次输入,在寄存器或共享内存中完成多个中间计算,最后将最终结果写回显存。这样,显存访问次数大幅减少,硬件利用率上升,训练速度显著提高。同时,因为不需要为每个中间结果分配显存,显存占用也能降低。

在Transformer中,常见的融合模式包括:

  • Conv + BatchNorm + ReLU 融合:在CNN中广泛应用。BatchNorm的归一化操作和ReLU的激活函数都是逐元素操作,完全可以嵌入到卷积的循环计算中,省去单独读写中间张量。对于Transformer中的MLP块,如果使用了BatchNorm(较少),原理相同。更常见的是Linear + LayerNorm/RMSNorm + Activation的融合。

  • QKV 投影融合:在自注意力中,Q、K、V的投影原本是三个独立的矩阵乘法:Q = X @ W_Q, K = X @ W_K, V = X @ W_V。可以将这三个权重矩阵拼接成一个更大的矩阵 W_QKV = [W_Q; W_K; W_V],然后执行一次矩阵乘法 X @ W_QKV,得到拼接后的结果。这减少了两个独立的矩阵乘法操作,降低了kernel启动次数和显存读写。

  • 自注意力分数计算 + Softmax + Dropout 融合:在FlashAttention出现前,计算 QK^T、softmax、Dropout(掩码操作)是多个独立步骤,每个步骤都会产生显存的写入和读取。将它们融合在一个kernel中完成,可以只读写一次Q和K,一次写出最终的注意力权重(或直接写V的加权结果)。

  • 残差连接 + LayerNorm/RMSNorm 融合:子层输出与原始输入相加后,立即进行归一化。这两个操作可以合并在一个kernel中,避免将相加的中间结果写回显存再读入归一化kernel。

  • MLP块的全连接 + 激活 + 全连接融合:第一个全连接层升维,经过GELU等激活,再降维。可以将第一个全连接和激活函数融合,将第二个全连接与第一个的激活直接在寄存器中衔接,减少显存读写。

  • AdamW优化器中的step融合:优化器更新参数时,涉及多个操作(如计算动量、更新参数、应用权重衰减),DeepSpeed等框架会将它们融合成一个或少数几个kernel,避免反复读写参数张量。

这些融合可以通过手写CUDA kernel、使用NVIDIA的CUTLASS库、或者依赖PyTorch 2.0的torch.compile自动完成。现代推理引擎(如TensorRT)更是将融合发挥到极致,以追求极致推理速度。

为什么将 Conv、BN 和 ReLU 融合成一个算子能加速训练?减少显存读写。

我们将这个经典融合作为案例,剖析其加速原理:

未融合的执行过程:

  1. 卷积操作:从显存读取输入特征图,执行矩阵乘加,将输出特征图写回显存。

  2. BatchNorm:从显存读取卷积的输出,计算均值、方差、归一化、缩放、偏移,将结果写回显存。

  3. ReLU:从显存读取BN的输出,逐元素应用max(0, x),将结果写回显存。

在这个过程中,特征图被完整地读写了三遍(每步一读一写)。特征图通常很大(例如[batch, 512, 56, 56]在FP16下约12MB)。每次读写都需要占用宝贵的显存带宽(A100的HBM带宽约2TB/s,但大模型训练中,带宽常被大量参数和梯度同步挤占)。

融合后的执行过程:

一个融合kernel被启动。它从显存读取输入特征图,在片上寄存器或共享内存中依次完成卷积、BN和ReLU的计算,将最终结果写回显存。整个过程特征图只被读写一次。

性能提升来源:

  • 减少显存带宽压力:读写次数降为原来的1/3,直接节省了大量时间。对于受带宽约束的操作(如卷积、归一化),这几乎是线性加速。

  • 消除kernel启动开销:每个独立kernel的启动都有固定延迟(几微秒到几十微秒)。融合后只需一次启动。

  • 提高数据局部性:中间数据直接保持在片上内存中,比反复从HBM读取快数十倍。

在Transformer训练中,这种融合对FFN层和归一化层尤其有效。虽然Transformer较少使用BatchNorm(多用LayerNorm),但LayerNorm、残差加法等同样可以进行类似融合。

在大模型训练中,如何将多个小矩阵乘法融合?

大模型(如LLaMA、GPT)中,除了巨大的矩阵乘(如[batch*seq, hidden] @ [hidden, 4*hidden]),还存在许多小矩阵乘。例如,在多头注意力中,每个头的Q、K、V投影通常被拆分为独立的矩阵乘法。这些“小矩阵乘”如果单独执行,效率很低,因为单个小矩阵乘无法充分利用GPU的Tensor Core和并行能力。

融合策略:

  • 拼接权重法(Batched GEMM):将多个小矩阵乘的权重矩阵拼接成一个大矩阵,输入也相应复制或广播,执行一次大矩阵乘,然后分割输出。例如,将Q、K、V的投影权重W_Q, W_K, W_V拼接为W_QKV,然后用X @ W_QKV一次性完成三个投影。这已经被广泛采用。

  • 使用torch.bmmtorch.matmul的批量模式:如果多个小矩阵乘具有相同的维度模式,可以构造成一个批量矩阵乘。例如,在GQA(分组查询注意力)中,KV头数少于Q头数,需要为每组Q头重复利用相同的K和V。可以通过扩展KV张量的维度,然后调用一次批量矩阵乘完成所有头的注意力计算。

  • 自定义CUDA kernel:专门为特定大小的小矩阵乘设计融合kernel,使用CUTLASS或手写CUDA,将多个小矩阵乘和后续的激活、缩放等操作融合在一起。例如,将自注意力的QK^T计算与softmax、与V的乘法融合为单个kernel(这正是FlashAttention的思路)。这种细粒度融合能带来最大的性能提升。

  • 框架自动融合(如torch.compile):PyTorch 2.0的Inductor后端可以自动分析计算图,将多个小矩阵乘融合为高效的CUDA kernel,无需手动编写。它会选择最优的tiling策略和内存访问模式。

  • FlashAttention的思路推广:将注意力机制中的两个矩阵乘(S = QK^TO = SV)以及softmax融合为单个kernel,完全避免了将中间注意力矩阵(形状[batch, heads, L, L])写入HBM,从而极大节省显存和带宽。

因此,对于大模型,优化小矩阵乘的核心是将它们“打包”成更大的矩阵乘,或者通过手写kernel将它们与周围的逐元素操作深度融合,避免中间结果的显存往返。

FlashAttention 是如何加速训练和节省显存的?写出其优化原理。

FlashAttention是Tri Dao等人提出的精确注意力算法,它不改变注意力计算的数学结果,却实现了显著加速和显存节省。其核心原理是IO感知的分块计算(tiling)。

问题背景:标准自注意力的计算需要实例化两个巨大的中间矩阵:注意力分数矩阵S = QK^T(形状[L, L])和softmax后的概率矩阵P(同样[L, L])。L为序列长度。这两个矩阵对显存带宽和容量都是灾难:它们必须写入HBM,然后再从HBM读回进行下一步计算(O = PV)。对于长序列,这个中间矩阵的大小可能超过模型参数本身。

FlashAttention的优化原理:

  • 将输入Q、K、V分块(Block):将序列长度维度切分成多个小块(Block),例如每块处理64或128个token。

  • 在每个Block内完成所有计算:加载一个Q块和一个K块到片上SRAM(共享内存/寄存器)。在SRAM内计算局部注意力分数,执行“在线softmax”(online softmax)更新,然后直接与对应的V块相乘,累加到输出块O。整个过程,完整的[L, L]矩阵从未被写入HBM。它只在SRAM中被临时计算和消费。

  • 在线softmax:传统softmax需要先求最大值、指数求和、最后归一化,必须获得整行数据。FlashAttention采用在线算法:每个块更新局部最大值和指数和,用它们修正之前块的结果,保证最终与全局softmax精确等价。

  • 优化内存访问:通过精心设计的数据加载模式,确保每次从HBM加载的Q、K、V块被充分利用,最大化SRAM内的计算密度。

加速和显存节省效果:

  • 显存节省:注意力矩阵的显存占用从O(L²)降为O(L)。以L=4096、头数32为例,注意力矩阵在FP16下原本需要约1GB显存,FlashAttention将其降为几MB的临时缓冲区。这使得训练更长的序列成为可能。

  • 速度提升:由于消除了对HBM的巨大读写量,计算效率提升,训练速度通常提升20%~40%,长序列下提升更为显著。

  • 实现:在PyTorch中,可通过torch.nn.functional.scaled_dot_product_attention启用FlashAttention,只需安装flash-attn包,并确保CUDA和硬件兼容。

FlashAttention 是否改变了注意力计算结果?精度损失吗?

FlashAttention的设计目标是数学上完全等价于标准注意力,因此理论上计算结果不应有任何改变。它只是改变了计算的执行顺序(分块、在线softmax),但所有运算(矩阵乘、softmax、求和)在数学上是严格相等的。

但是,实践中存在微小的精度差异,主要来源于:

  • 浮点运算的非结合性:分块计算改变了求和的顺序。例如,全局softmax是先算完整行的最大值和指数和,而在线softmax是逐块更新。在FP16/BF16等有限精度下,不同的求和顺序可能产生不同的舍入误差,导致结果有微小偏差(通常在1e-5量级)。这与标准注意力在GPU上由于并行归约顺序不固定导致的误差类似。

  • 硬件架构差异:FlashAttention可能会使用Tensor Core或其他特定硬件指令,与标准实现的数值路径略有不同。

  • 随机Dropout:如果使用了注意力Dropout,由于需要生成随机mask,即使种子相同,不同算法实现可能产生不同mask(通常框架会保证一致性)。

总体而言,这种精度差异对训练收敛和最终模型质量没有可观测的影响。FlashAttention已在GPT-3、LLaMA等模型的训练和推理中被广泛验证。如果出于严格复现的考虑,可以通过固定随机种子和使用确定性算法来最小化差异,但无此必要。

在 PyTorch 中,如何启用 FlashAttention?torch.nn.functional.scaled_dot_product_attention。

PyTorch从2.0版本开始,提供了统一的注意力接口torch.nn.functional.scaled_dot_product_attention(简称SDPA)。该接口会自动根据输入张量的形状、CUDA能力、是否可用FlashAttention等条件,选择最优的底层实现。

启用步骤:

  1. 安装flash-attn库:pip install flash-attn --no-build-isolation。注意,FlashAttention对CUDA版本、GPU架构(Ampere以上,如A100、H100、RTX 3090/4090)有要求。

  2. 在代码中直接调用SDPA:

import torch
import torch.nn.functional as F

# Q, K, V 形状为 (batch, heads, seq_len, head_dim)
attn_output = F.scaled_dot_product_attention(Q, K, V, attn_mask=mask, dropout_p=0.0, is_causal=True)
  1. PyTorch会自动使用FlashAttention,前提是满足条件(如无自定义attn_mask的某些模式,或不支持某些参数组合)。你可以通过环境变量TORCH_LOGS="+graph_breaks"torch.backends.cuda.sdp_kernel来查看或控制具体实现。

  2. 如果想强制使用FlashAttention,可设置:

torch.backends.cuda.enable_flash_sdp(True)
  1. 对于HuggingFace模型,很多已集成了use_flash_attention_2=True参数,只需在加载模型时指定即可。

SDPA的优势:它还能自动切换到其他高效实现,如xFormers的memory efficient attention,为不同的硬件和输入配置选择最合适的后端。这使得模型代码无需修改即可享受加速。

除了 FlashAttention,还有哪些高效的注意力实现?(如 xFormers 的 memory efficient attention)

  • xFormers (Meta出品):提供memory_efficient_attention,使用类似的tiling技巧,但支持更多类型的注意力偏置(如相对位置偏置、ALiBi等)。它在非Ampere架构上也有良好的性能,且与FlashAttention兼容,常作为SDPA的后端之一。

  • PagedAttention (vLLM中):专门针对推理场景,将KV缓存分页管理,允许在非连续显存中存储,解决KV缓存碎片化问题,极大提升推理吞吐。

  • Ring Attention (分布式):将序列长度维度切分到多个GPU上,通过环形通信进行注意力计算,使训练上下文窗口扩展到数万甚至十万token。结合FlashAttention,可实现高效的分布式长序列训练。

  • Sparse Attention / Dilated Sliding Window:如Longformer、BigBird,通过稀疏或膨胀窗口限制每个token只关注部分token,复杂度从O(L²)降为O(L·W)(W为窗口大小)。适合超长序列。

  • Multi-Query / Grouped-Query Attention (MQA/GQA):在模型结构层面减少KV缓存,间接加速推理。FlashAttention等均支持这些变体。

  • HyperAttention:利用低秩近似和哈希技术,进一步加速长序列注意力。

  • StreamingLLM:通过保留少量“attention sink” token,实现无限长度输入的高效推理。

这些实现通常可组合使用,例如在分布式长序列训练中,使用Ring Attention结合FlashAttention,在节点间进行通信的同时,节点内利用FlashAttention进行快速计算。

什么是“激活值重计算” (Activation Recomputation)?与梯度检查点是一回事吗?

激活值重计算与梯度检查点 (Gradient Checkpointing) 实质上指的是同一技术,只是名称上的侧重点略有不同。

原理:在标准训练中,前向传播产生的中间激活值需要全部保存在显存中,因为反向传播需要它们来计算梯度。当模型层数很深或序列很长时,这些激活值会占用巨大的显存,成为训练的主要瓶颈。

激活值重计算的核心思想是:用计算换显存。在前向传播时,选择性地丢弃某些层的中间激活值(不保存)。反向传播需要这些激活值时,从最近的保存点(检查点)开始,重新执行该段的前向计算,临时恢复被丢弃的激活值,用于梯度计算,随后立即释放。

实现:通常将一个Transformer层或一段连续的层设为一个检查点段。前向时只保存该段的输入张量(检查点),丢弃段内所有中间激活。反向时,用保存的输入重新进行前向传播,得到所需中间激活,然后计算该段梯度,释放临时激活。

与梯度检查点的异同:

  • 完全相同。PyTorch中通过torch.utils.checkpoint.checkpoint实现,也称为“activation checkpointing”。

  • 该技术不改变模型结构和训练算法,仅改变内存管理策略。

  • 代价是增加了一次额外的前向计算(约33%的额外计算量),但能节省50%~70%的激活值显存。在大模型长序列训练中,这是必不可少的技术。

因此,激活值重计算就是梯度检查点,两者在工业界和学术界通用,只是描述角度不同:一个强调“重计算激活值”,另一个强调“设置检查点”。

梯度检查点如何用20%的计算换取50%的显存?

在标准的训练过程中,前向传播会产生大量中间激活值(比如每一层的输出、注意力矩阵、FFN中间结果等),这些激活值在反向传播时需要被消费来计算梯度,因此必须全部保存在显存中。对于一个大模型,激活值显存可以轻易超过模型参数本身,成为限制batch size和序列长度的瓶颈。

梯度检查点的核心思想很简单:不要把所有激活值都存着,只存一小部分“检查点”,反向传播时再从检查点重新计算被丢弃的激活值。 这相当于把显存压力转移到了计算上。

具体实现中,我们通常以每个Transformer层为单元设置检查点。前向传播时,这一层内部的Q、K、V、注意力矩阵、FFN中间输出等全部不保存,只保留该层的输入张量(即上一层的输出)。当反向传播推进到这一层时,利用保存的输入重新执行一次前向传播,临时生成所需的中间激活,用完立即释放。

为什么是“20%计算换50%显存”?

  • 对于一个典型的Transformer层,前向计算量约占整个训练步(前向+反向)的1/3。因为反向传播的计算量大约是前向的两倍(需要计算权重梯度和输入梯度)。梯度检查点让被标记的层多执行一次前向,因此额外的计算开销大约是 1/3 * (检查点层数/总层数)。如果所有层都开启检查点,额外计算量约33%。但很多时候我们只对部分层开启(如每隔一层),所以额外开销通常在15%-25%之间,四舍五入就是“20%”。

  • 显存节省方面,每一层的中间激活通常占据总激活显存的很大比例。以一个标准的Transformer层为例,如果不存注意力矩阵,该层的激活量大约为 batch * seq_len * hidden_dim * (几十倍)。开启检查点后,这一整层的中间激活全部不存,只保留一个输入张量。对于深层模型,激活总量可减少50%甚至70%。所以“50%”是一个保守但具有代表性的数字。

实际效果举例:训练一个7B模型,序列长度2048,batch size=1。不使用检查点时,激活显存大约需要20-25GB。开启全部层的检查点后,激活显存降至约8-10GB,节省了约60%。同时单步训练时间从0.5秒增加到0.6秒,增加约20%。这就是“20%计算换50%显存”的由来。

梯度检查点通常设置在 Transformer 的哪一层?为什么?

通常设置在整个Transformer层的粒度,即以一个完整的Transformer块(包含自注意力和FFN,或Pre-Norm下的子层)作为检查点单元。原因有三:

最大显存节省:一个Transformer层内部产生的中间激活是整个网络激活值的绝对主体。注意力矩阵([batch, heads, L, L])和FFN的中间输出([batch, L, 4*hidden])占用了最大头。将整层设为检查点,这些大张量全部丢弃,节省效果最显著。

重计算粒度适中:如果以更细的粒度(比如注意力层和FFN层分别设检查点),虽然可以更灵活,但实现复杂,而且会增加重计算的次数(额外前向传播会增多)。如果以更粗的粒度(比如几层一起),则重计算时需要重跑好几层,额外计算量更大。单层检查点在计算开销和显存节省之间取得了最佳平衡。

实现简洁:现代框架(如PyTorch的torch.utils.checkpoint)可以非常方便地将一个nn.Module标记为检查点,只需在forward调用时包裹即可。以层为单位最自然。

例外:有时候我们会对嵌入层(Embedding)和输出层(LM Head)单独处理,因为这些层的激活值通常不大,而且参数量巨大,重计算代价高,通常不作为检查点。

PyTorch 中如何实现梯度检查点?torch.utils.checkpoint。

在PyTorch中,使用torch.utils.checkpoint.checkpoint函数来包裹一个模块的前向传播。基本用法:

import torch
from torch.utils.checkpoint import checkpoint

class TransformerLayer(nn.Module):
    def __init__(self, ...):
        self.self_attn = MultiHeadAttention(...)
        self.ffn = FeedForward(...)
        self.norm1 = nn.LayerNorm(hidden_dim)
        self.norm2 = nn.LayerNorm(hidden_dim)

    def forward(self, x):
        # 使用检查点包裹自注意力块
        x = x + checkpoint(self.self_attn, self.norm1(x), use_reentrant=False)
        # 使用检查点包裹FFN块
        x = x + checkpoint(self.ffn, self.norm2(x), use_reentrant=False)
        return x

use_reentrant=False是推荐设置,它使用非重入版本,避免了在多线程环境下的一些潜在问题,且支持更多的PyTorch特性。

也可以对整个Transformer层的forward做一次包裹,但这会重计算整个层,而非细粒度的子层。实际中可以根据需要决定粒度。

注意:检查点要求输入张量的requires_grad=True,且模型在训练模式下。被包裹的函数不应该有随机操作(如Dropout),因为重计算时会重新随机化导致不一致。PyTorch会尝试在重计算时恢复相同的随机状态,但保险起见,通常将Dropout放在检查点之外,或者在包裹的函数内部手动设置seed。

使用梯度检查点后,训练总时间增加多少?如何测量?

增加的时间取决于开启检查点的层数占比。理论上,若所有层都开启,额外计算量约为总计算量的30%-50%(因为前向传播占比约1/3,反向约2/3,重算一次前向意味着增加1/3的计算量)。但由于GPU并行性和通信等因素,实际时间增加可能略低或略高。

测量方法:

  1. 在相同的硬件和数据上,分别运行开启和不开启检查点的训练,各跑几百步,确保预热完毕。

  2. 记录每秒处理的样本数(throughput)或单步平均耗时。

  3. 计算时间增加百分比:(T_checkpointed - T_baseline) / T_baseline * 100%

实际数据:以LLaMA-7B为例,在A100单卡上,不开启检查点时,单步约0.45秒;开启所有层检查点后,单步约0.58秒,增加约29%。如果只开启一半层,则增加约15%。

注意:时间增加还与序列长度有关。长序列下注意力计算占比更大,检查点重计算注意力矩阵的时间也就更多,因此时间增幅可能略高。

为什么对 MLP 层做检查点比对注意力层更划算?

这里的“划算”指的是用更少的额外计算换取更大的显存节省。

  • 显存节省:MLP层的中间激活尺寸通常是 [batch, L, 4*hidden],而注意力层的中间激活除了QKV([batch, L, hidden]),还有注意力矩阵([batch, heads, L, L])。对于长序列(L很大),注意力矩阵显存占主导。但对于短序列或使用了FlashAttention(不保存注意力矩阵),MLP中间激活的显存占比可能更大。所以在不同场景下,两者都可能成为显存大户。

  • 计算开销:MLP层的计算主要是两个巨大的矩阵乘法(升维和降维),其计算量(FLOPs)通常远大于注意力层(在序列较短时)。重计算MLP层的代价更高。而注意力层除了QKV投影外,还有QK^TPV的矩阵乘,在长序列时计算量也很大。

因此,更划算的做法通常是优先对MLP层做检查点,或者对注意力层和MLP层都做检查点。单独对注意力层做检查点的情况较少,因为注意力矩阵可以通过FlashAttention避免存储,此时注意力层的激活值已经很小,无需检查点。在FlashAttention普及前,注意力矩阵是最大瓶颈,因此那时也会对注意力层做检查点。现在有了FlashAttention,很多框架默认只对FFN做检查点,因为FFN的中间激活依然很大(如LLaMA的SwiGLU FFN),而且重计算FFN的代价相对于其节省的显存来说是可接受的。

总结:在现代大模型训练中,我们通常结合FlashAttention和FFN的梯度检查点,将注意力矩阵的存储消除,并将FFN的中间激活以较小的计算代价压缩。这样在不显著增加计算的前提下,最大化了显存节省。

什么是“选择性的梯度检查点”?只重计算一部分激活。

选择性的梯度检查点(Selective Checkpointing)是指并非对所有层或所有激活值都进行重计算,而是只选择那些显存占用大、重计算代价小的激活值进行丢弃和重计算。

实现方式:

  • 按层选择:比如只对FFN层开启检查点,注意力层保留全部激活(因为注意力矩阵已被FlashAttention优化)。

  • 按激活张量选择:在自定义的autograd Function中,手动控制哪些中间张量被保存,哪些被丢弃。例如,可以保存QKV但丢弃注意力矩阵,或保存部分归一化层的输出。

  • 利用框架特性:DeepSpeed和Megatron都支持不同粒度的检查点。DeepSpeed可以配置partition_activationscontiguous_checkpointing等选项。Megatron的recompute模块允许指定recompute_method='uniform'(每层都重算)或'block'(按块重算)。

好处:最大化显存节省的同时最小化计算开销。对于一些重计算代价高的操作(如大矩阵乘),可以选择保留其激活;对于重计算代价低的操作(如逐元素激活函数),可以放胆丢弃。

实践:在LLaMA等模型中,通常对每个Transformer层的整个forward开启检查点,而FlashAttention内部已经自动处理了注意力矩阵的不存储。因此,最终的重计算主要落在QKV投影和FFN层上。

编译器加速:torch.compile 是如何优化训练图的?

torch.compile是PyTorch 2.0引入的即时(JIT)编译器,它使用TorchDynamo捕获PyTorch的字节码,将其转换为FX图,然后交给后端(如Inductor)进行优化和代码生成。

优化流程:

  1. 图捕获:TorchDynamo在运行时拦截Python的frame,将PyTorch操作记录为一张计算图。

  2. 图优化:Inductor后端对图进行分析,执行一系列优化pass,包括:

  3. 算子融合:将相邻的逐元素操作(如激活函数、归一化、加法)融合成单个kernel,减少显存读写。
  4. 水平融合:将多个相同形状的矩阵乘(如QKV投影)合并为一次更大的矩阵乘。
  5. 内存规划:优化中间缓冲区的分配和复用,减少显存碎片和浪费。
  6. 自动混合精度:根据硬件特性,自动选择最优的数据类型。

  7. 代码生成:生成高效的CUDA C++代码或Triton kernel,利用GPU的硬件特性(如Tensor Core、共享内存)。生成的代码是高度特化的,针对输入张量的具体大小和步幅进行了优化,避免了通用kernel中的动态判断开销。

  8. 缓存:编译结果被缓存,相同图结构和输入大小可以直接复用,避免重复编译。

因此,torch.compile可以将PyTorch的动态图“静态化”,在底层应用大量手工优化级别的加速,同时保持了PyTorch的易用性。

torch.compile 的 “mode” 参数有哪些?"default", "reduce-overhead", "max-autotune" 的区别。

torch.compilemode参数控制编译优化的激进程度:

  • "default":平衡模式。编译时应用大部分优化,但不会花费太多时间在自动调优上。通常能获得不错的速度提升,编译时间较短。适合大多数训练场景。

  • "reduce-overhead":侧重于减少框架开销。它会尝试将更多操作融合,减少Python-CUDA交互次数。可能会生成更激进的融合kernel,编译时间稍长,但训练循环的延迟更低,尤其适合小模型或推理场景。

  • "max-autotune":极致性能模式。编译器会花费大量时间对不同的实现方案进行benchmark(如矩阵乘的tiling策略、线程块大小),选择在具体硬件上最快的那一个。编译时间最长(可能数小时),但能获得最优性能。适用于对吞吐有极致要求的生产环境,或模型固定后的最终优化。

实际使用建议:训练大模型时,先用"default"模式快速验证。如果追求长期训练的极致效率,可以切换到"max-autotune"并让它在夜间自动完成编译,然后加载缓存进行正式训练。

在使用 torch.compile 时,为什么第一次迭代特别慢?(编译开销)

第一次迭代慢是因为编译器在进行图捕获、分析和代码生成。这个过程包括:

  • TorchDynamo捕获字节码,构建计算图。

  • Inductor对图进行优化(融合、内存规划)。

  • 调用Triton或CUDA编译器将优化后的图编译成可执行的GPU代码。

  • 编译好的kernel需要加载到GPU并缓存。

对于大模型,图规模庞大,编译时间可能长达几分钟甚至数十分钟。而且由于动态形状,不同的输入形状可能触发重新编译。

如何缓解:

  • 预热(warmup):在正式训练前,用真实输入大小跑几个空batch(只前向不反向),触发编译并缓存。之后正式训练就不会有编译开销了。

  • 缓存持久化:编译好的代码默认缓存到~/.cache/torch/。下次使用相同的图和大小时,直接从缓存加载,几乎无延迟。

  • 固定输入形状:尽量使用固定大小的输入(如固定的序列长度),避免频繁的形状变化导致重新编译。

如何通过 torch.compile 结合 Tensor Cores 加速?

Tensor Cores是NVIDIA GPU上专为矩阵乘累加设计的硬件单元,能提供比普通CUDA Core高数倍的吞吐。要利用Tensor Cores,需要满足特定的矩阵维度和数据布局要求(如FP16/BF16,且维度为8或16的倍数)。

torch.compile的Inductor后端可以自动生成利用Tensor Cores的代码。它会:

  1. 自动检测输入张量的数据精度和形状,如果符合Tensor Cores的使用条件,就会生成对应的mma(matrix multiply-accumulate)指令。

  2. 在融合kernel中,将矩阵乘法部分映射到Tensor Cores,逐元素操作(如激活函数)映射到普通CUDA Cores,两者在同一个kernel中协同工作,避免数据反复进出HBM。

  3. 通过自动调优,选择最佳的tiling和线程块配置,最大化Tensor Cores的利用率。

用户无需手动干预。只需确保模型使用FP16或BF16混合精度训练,模型结构(如注意力头数、隐藏维度)尽量为8或16的倍数。torch.compile就会自动生成Tensor Cores优化代码。如果希望在日志中查看是否使用了Tensor Cores,可以设置环境变量TORCH_LOGS="+output_code"查看生成的Triton/CUDA代码,或使用Nsight Compute进行profiling。

实践:我在训练ViT和GPT模型时,使用torch.compile前后,在A100上MatMul吞吐提升了约20%-40%,部分得益于Tensor Cores的有效利用。需要说明的是,即使不使用torch.compile,PyTorch默认的cuBLAS后端也会调用Tensor Cores,但torch.compile通过融合和减少显存访问,能让Tensor Cores的“有效利用率”更高,避免因数据搬运导致的停顿。

综上所述,梯度检查点和torch.compile都是大模型训练中不可或缺的性能优化工具。前者用计算换显存,后者用编译优化换速度,搭配使用往往能获得1+1>2的效果。

什么是“动态图”与“静态图”?torch.compile 如何将动态图转为静态图?

动态图(Dynamic Graph) 是PyTorch默认的执行模式。每次前向传播时,框架都会即时构建一张计算图,执行完就丢弃。这给了我们极大的灵活性:你可以在forward里写if-else、for循环、动态形状变化等,就像写普通Python代码一样自然。但代价是运行时开销大——每次都要重新建图、调度算子、调用CUDA kernel,这些“框架开销”在大规模训练中累积起来相当可观。举个例子,如果模型有几百个小算子,每次前向传播都要为每个算子单独启动一个CUDA kernel,每个kernel启动都有几微秒到几十微秒的延迟,加起来就可能达到毫秒级,对于计算量本身只有几毫秒的操作来说,框架开销占比甚至超过50%。

静态图(Static Graph) 则是在运行前先定义好完整的计算图,编译优化后再执行。TensorFlow 1.x就是这个路数。优点是执行效率高,可以做全局图优化(算子融合、内存规划、常量折叠等),缺点是不灵活,调试困难。比如你想在forward里根据某个中间结果的大小动态决定后续操作,在静态图里就很别扭。

torch.compile 的魔法就是在这两者之间架桥。它使用TorchDynamo在运行时捕获Python字节码。具体来说,TorchDynamo会拦截Python解释器的frame执行,把PyTorch的算子调用记录成一张FX计算图。这个过程对用户透明——你依然可以随意写动态逻辑,但TorchDynamo会智能地识别出哪些部分是可以“静态化”的。

捕获到的图会被交给Inductor后端。Inductor对图进行大量优化:比如把相邻的逐元素操作(GELU、Dropout、残差加法)融合成单个CUDA kernel,消除中间张量的显存读写;比如把Q、K、V三个投影矩阵拼成一个大矩阵,做一次大矩阵乘而不是三次小矩阵乘;再比如自动为矩阵乘法选择最优的tiling策略和线程块大小。这些优化之后,Inductor生成高效的Triton或CUDA代码,编译结果被缓存,下次直接加载。

对于图中无法静态化的部分(比如依赖Python控制流的动态分支),torch.compile会将其保留为“图断点”(graph break),恢复为eager模式执行。这样既保留了动态图的灵活性,又在大部分计算密集区域获得了静态图的性能。最终效果是,我们无需修改任何代码,就能让训练速度提升20%-50%。

大规模训练中,数据加载经常成为瓶颈,如何优化?

数据加载瓶颈的表现是:GPU利用率忽高忽低,经常掉到50%以下,而CPU利用率却很高。这说明GPU在等数据,CPU在拼命准备数据,但供不应求。这种情况在大模型训练中极为常见,因为模型计算量大、速度快,而数据预处理(尤其是图像解码、文本分词等)却很吃CPU。

优化方向:

  • 增加并行度:通过num_workers参数开启多进程数据加载。每个worker独立处理一个batch的数据,主进程从队列中取。一般设为4-16,但并非越大越好——太多worker会争抢CPU和内存,每个worker都会fork一份主进程的内存,如果主进程cache了数据集,内存会翻很多倍。建议从CPU核心数的一半开始调试,同时用htop监控CPU和内存使用。

  • 避免Python GIL:使用多进程而非多线程,因为Python的全局解释器锁会限制线程的并行能力。num_workers就是多进程模式。

  • 减少IO阻塞:将数据存放在高速NVMe SSD上,或使用分布式文件系统(如Lustre、GPFS)增加聚合带宽。避免NFS等低并发存储。在多机训练时,最好每个节点本地都有数据副本,或者使用能够支持高并发读的并行文件系统。

  • 预处理离线化:将分词、截断、拼接等操作提前处理好,保存为二进制文件(如numpy数组、Arrow格式),训练时直接加载数值序列,省去在线处理的开销。这是最彻底的优化方式,尤其适用于文本大模型。

  • 使用高效数据格式:避免逐个小文件读取(元数据压力大),使用打包格式(如TFRecord、WebDataset的tar包)。对于图像数据,将成千上万张小图打包成一个tar文件可以大幅减少随机IO。

  • CPU-GPU数据传输优化:开启pin_memory=True,使用页锁定内存加速主机到设备的数据拷贝;配合non_blocking=True让拷贝与计算重叠。页锁定内存不会被操作系统换出,GPU的DMA引擎可以直接高速访问。

  • GPU直接数据加载:使用NVIDIA DALI,将数据预处理(解码、增强)从CPU卸载到GPU,消除CPU瓶颈。这在图像、视频等需要复杂解码的场景中效果尤其显著。

  • 预取与缓存:使用内存缓存常用数据,减少磁盘IO。对于较小的数据集,可以直接把整个处理好的数据集放到内存里。

实践中,我通常会先跑一个profiler,用torch.utils.bottlenecknsys看看哪个环节耗时最长。如果发现dataloadergetitem耗时高,就优化预处理或增加worker;如果磁盘IO是瓶颈,就上SSD或分布式存储;如果CPU到GPU的传输慢,就开pin_memorynon_blocking

使用 DALI (NVIDIA Data Loading Library) 进行 GPU 直接数据加载的好处。

DALI是NVIDIA专门为深度学习训练打造的数据加载库,它把数据预处理的流水线完全放在GPU上执行,从而彻底消除CPU瓶颈。与传统方式(CPU解码图像、resize、归一化,再传到GPU)不同,DALI直接在GPU上完成解码、缩放、裁剪、颜色抖动等一系列操作,处理好的张量直接用于模型前向传播,无需再从CPU拷贝。

核心好处:

  • CPU负载大幅降低:复杂的图像/音频解码和增强不再占用CPU,CPU只需负责从磁盘读取原始字节流并喂给GPU。这对于CPU核心数有限或者需要同时跑多个训练任务的服务器尤为重要。

  • 端到端GPU流水线:数据增强和模型计算都在GPU上,可以利用CUDA Stream实现重叠——当前batch训练时,下一个batch的预处理已在后台完成。这种“生产者-消费者”模式让GPU几乎感受不到数据加载的延迟。

  • 显存优化:DALI处理完的Tensor直接留在GPU显存中,避免了CPU-GPU间的多次往返拷贝。对于高分辨率图像,每次拷贝都是几十上百MB的数据,省去这些拷贝对训练吞吐提升明显。

  • 支持多种数据格式:不仅图像,还支持视频、音频、点云等,非常适合多模态训练。比如训练一个视频理解模型,DALI可以直接在GPU上解码视频帧,效率远高于CPU解码。

  • 易于扩展:与PyTorch的DataLoader无缝集成,可以替代原有的Datasettransform。DALI提供了DALIGenericIterator,使用方式和PyTorch的DataLoader类似。

使用DALI后,图像模型的训练吞吐往往能提升1.5-2倍,GPU利用率稳定在95%以上。对于大模型的多模态训练(如CLIP、Stable Diffusion),DALI几乎是标配。不过需要注意,DALI的学习曲线稍陡,配置pipeline需要一些调试经验,而且某些复杂的自定义增强逻辑不太容易用DALI表达。

什么是“数据预取” (Data Prefetching)?num_workers 和 prefetch_factor 的设置。

数据预取是指在GPU处理当前batch的同时,CPU在后台提前准备好后续几个batch的数据。这样GPU永远不会因为等数据而空闲,形成“计算-数据准备”的流水线。

在PyTorch的DataLoader中,num_workers控制预取进程的数量,prefetch_factor控制每个worker预先准备的batch数量。总预取队列长度 = num_workers * prefetch_factor

  • num_workers:每个worker是一个独立的Python进程,并行执行getitem。通常设为4-8,对于重CPU任务(如图像解码)可设到16甚至更高。但worker太多会消耗大量CPU内存,因为每个worker都会fork主进程的内存空间(如果cache了大数据集,内存占用会翻倍)。建议用htop监控CPU内存和利用率,找到最优值。一个经验法则是:观察一个step中GPU空闲时间,逐步增加worker数直到GPU不再空闲。

  • prefetch_factor:每个worker提前加载的batch数。默认为2。增大它可以让队列更长,更能应对数据加载的瞬时波动,但会占用更多CPU内存。如果数据加载很稳定,设为2就够;如果数据预处理耗时波动大(比如不同样本的处理时间差异大),可以设到4-8。

  • persistent_workers:设为True可以保持worker进程存活,避免每个epoch重新fork的开销。但注意,如果数据集的epoch间shuffle逻辑有问题,可能会导致重复数据。

最佳实践:先设num_workers=4, prefetch_factor=2, persistent_workers=True跑几个step,用nvidia-smi观察GPU利用率是否稳定且高。如果GPU经常空闲,逐步增加worker数;如果CPU内存告急,减少worker或使用轻量级预处理。还可以在DataLoader中使用pin_memory=Truepin_memory_device来加速数据传输。

内存映射 (Memory Mapping) 在数据加载中的应用,减少 CPU 内存占用。

内存映射(mmap)是一种让应用像访问内存一样直接访问磁盘文件内容的技术,由操作系统负责数据的按需加载。在训练大模型时,我们通常会将原始文本语料提前分词并转换成整数 token 序列,然后保存为巨大的二进制文件。对于 TB 级别的数据集,直接把这些文件全部加载进 CPU 内存是不现实的。

工作原理:mmap 将一个文件映射到进程的虚拟地址空间,但并不会立即把文件内容拷贝到物理内存。当程序访问某个地址时,操作系统发现该页不在物理内存中,会产生一个缺页中断,然后自动从磁盘读取对应的数据块。这实现了“按需加载”。而且操作系统会利用页缓存(page cache)进行智能预读和缓存,访问过的数据会留在内存中,后续访问就变成了内存访问。

在大模型训练中,我们对 token 序列的访问通常是顺序的(从头到尾逐个 epoch),这种访问模式非常适合 mmap。我们只需在创建数据集时,使用 np.memmap 或 PyTorch 的 torch.UntypedStorage.from_file 来映射二进制文件,然后像操作普通数组一样切片读取。操作系统会自动管理内存,整个训练过程中,CPU 内存占用只维持在几百 MB 到几 GB 的水平,而不会撑爆内存。

减少 CPU 内存占用的关键:

  • mmap 映射的文件不占用物理内存,只占用虚拟地址空间。当多个 worker 进程同时访问同一文件时,操作系统会共享页缓存,进一步节省内存。

  • 可以配合 madvise 系统调用,提示内核预取策略或指示即将访问的区域,让 IO 更加高效。例如,使用 MADV_SEQUENTIAL 告诉内核我们是顺序访问,内核会主动预读后续数据并释放已读过的页。

  • 如果是分布式训练,每个节点都可以映射存储在本地的数据文件。如果数据在分布式文件系统上,mmap 可能因网络延迟而导致性能下降,此时更适合将数据先拉到本地 NVMe 再映射。

经验教训:有一次我们发现训练过程中 CPU 内存持续上涨,最后 OOM,排查发现是因为 DataLoader 在每个 epoch 都重新 np.memmap,而旧映射没有及时释放导致内存泄漏。正确的做法是在 init 中一次性创建映射,在 getitem 中只读取切片。另外,不要对 mmap 对象进行 pickle 序列化,在 num_workers>0 时,应在每个 worker 的 worker_init_fn 中独立打开文件映射,避免跨进程共享带来的锁冲突。

如何压缩训练数据以减少 I/O?(如使用 Arrow 格式)

训练数据的 I/O 瓶颈往往是整个训练流水线的最大短板。压缩数据可以减少磁盘读取量和网络传输时间,但压缩本身也会消耗 CPU,需要在解压开销和 I/O 收益之间权衡。

Arrow 格式:Apache Arrow 是一种列式存储格式,天然支持零拷贝读取和高效压缩。与传统的 row-based 格式(如 CSV、JSON)不同,Arrow 按列组织数据,这使得:

  • 读取某一列时无需加载整行,适合需要部分字段的场景。

  • 同列数据类型相同,可以高效压缩(如字典编码、游程编码、位图压缩等)。

  • 可以直接在内存中进行操作,无需反序列化。PyTorch 可以直接从 Arrow 的 buffer 创建 Tensor,实现零拷贝。

对于文本大模型,我们可以把分词后的 token 序列存储为 Arrow 的 List<Int64> 列,配合 Snappy 或 LZ4 等轻量压缩算法。这些算法的特点是压缩率适中(通常2-4倍),但解压速度极快,不会成为 CPU 瓶颈。相比之下,Gzip 压缩率更高但解压太慢,不适合训练场景。

其他压缩策略:

  • 二进制打包:将多个短样本拼接成固定长度的序列,用特殊分隔符隔开,然后以二进制格式存储。这避免了大量小文件,减少 IO 次数。

  • 增量编码:对于有序的 token 序列,可以存储相邻 token 的差值(delta),然后用变长整数编码,大幅减少文件体积。

  • WebDataset/tar 包:将海量小文件打包成 tar 文件,减少随机读取和元数据压力。读取时顺序解包,配合内存缓存,效率很高。

在实践中,我通常会把预处理好的 token 序列用 Arrow 的 Feather 格式(Arrow 的文件存储格式)保存,配合 LZ4 压缩。读取速度接近原生 numpy,但文件体积只有原来的 1/3 左右。这样在多机训练时,可以更快地将数据从存储拉取到各节点。

训练时,使用“混合精度”不仅省显存,还能加速,为什么?

混合精度训练是指在前向和反向传播中使用半精度浮点数(FP16或BF16),而关键权重副本和部分计算保留全精度(FP32)。它既节省了显存,又加速了计算。

加速原理:

  • Tensor Core 专门优化:现代 NVIDIA GPU(Volta及以后)都有专门的 Tensor Core 硬件单元,专为 FP16/BF16 矩阵乘加操作设计。它的吞吐量是同等 FP32 操作的数倍甚至十多倍。混合精度训练将大型矩阵乘法(如 QKV 投影、FFN 层、注意力计算)都放到 Tensor Core 上执行,极大地提升了计算速度。

  • 显存带宽减半:FP16/BF16 张量占用的显存是 FP32 的一半。这意味着在同样的显存带宽下,可以传输两倍的数据,有效缓解了“内存墙”瓶颈。对于归一化、激活函数等带宽受限的操作,速度自然提升。

  • 算子融合更高效:在低精度下,GPU 的寄存器可以容纳更多的数据,有利于进行更激进的算子融合,进一步减少显存读写次数。

为什么省显存:模型参数、梯度、激活值都以 FP16/BF16 存储,显存占用直接减半。虽然还需要一份 FP32 的主权重用于精确更新,但总显存仍然大幅降低。这使得我们可以在同样的硬件上训练更大的模型或使用更大的 batch size。

Tensor Cores 需要什么条件才能启用?(如矩阵尺寸对齐)

Tensor Cores 是 NVIDIA GPU 中的专用硬件单元,专为加速矩阵乘加运算(D=A*B+C)而设计。要启用 Tensor Cores,需要满足以下条件:

精度要求:

  • 输入矩阵必须为 FP16、BF16、TF32、INT8 或 INT4 等特定格式。FP32 的矩阵乘法不会使用 Tensor Cores(除非启用 TF32 模式,将 FP32 截断为 TF32 用于计算,结果仍为 FP32)。

矩阵尺寸对齐:

  • 对于 FP16/BF16,矩阵的维度(m, n, k)需要对齐到 8 的倍数。这是因为 Tensor Core 每次处理一个固定大小的矩阵块(如 16×16×16)。如果不对齐,cuBLAS 会自动 padding 或回退到普通 CUDA Core 实现,导致性能大幅下降。

  • 在实际模型设计中,隐藏维度(hidden_dim)通常设为 512、768、1024、4096 等 8 的倍数;注意力头数也设为 8 的倍数。这是专门为了最大程度利用 Tensor Cores。

其他条件:

  • 使用 cuBLAS 或 cuDNN 等库调用,并在代码中启用 Tensor Cores(PyTorch 默认已开启)。

  • 对于某些操作,需要显式设置 torch.backends.cuda.matmul.allow_tf32 = True(在 Ampere 及更新架构上)来允许使用 TF32。

为什么尺寸对齐很重要:如果矩阵尺寸不对齐,比如隐藏维度是 1025,那么 cuBLAS 无法高效地将其映射到 Tensor Core 的固定块上,只能回退到普通 CUDA Core,速度会慢 3-5 倍。因此,在设计模型结构时,始终将关键维度设为 8 或 16 的倍数。

为什么训练中推荐 batch size 对齐到 8 的倍数?与 Tensor Core 有关。

如上一问所述,Tensor Cores 的矩阵乘加操作要求矩阵维度对齐到 8(FP16/BF16)的倍数。当输入矩阵的 m 维度(对应 batch_size * seq_len)不能整除 8 时,硬件无法高效利用 Tensor Core 的完整计算能力,性能会显著下降。

在实际训练中,矩阵乘的 m 维度通常是 batch_size * seq_len 的乘积。为了让这个乘积保持 8 的倍数,有两种做法:

  • 让 batch_size 对齐到 8:这比较简单,直接设置 batch_size=8, 16, 32 等即可。

  • 让每个 GPU 处理的 token 总数(micro_batch_size * seq_len)对齐到 8:如果序列长度是固定的(如 2048),那么 micro_batch_size 设为 1、2、4、8 都能保证乘积是 8 的倍数。如果序列长度变化,或者使用动态 padding,就需要在数据加载时进行填充,使每个 batch 的 token 总数对齐到 8。

因此,batch size 对齐 8 并不是强制要求,但它是确保 Tensor Core 高效利用的最简单方法。现代大模型训练几乎都会将 micro_batch_size 设为 1 或 2,此时需要关注的是 seq_len 是否对齐到 8。如果对齐了,即使 batch size 很小,Tensor Core 也能高效工作。

什么是“Auto Mixed Precision” (AMP)?PyTorch 中如何使用?

AMP(自动混合精度)是 PyTorch 提供的自动化混合精度训练方案。它自动决定哪些操作应使用 FP16/BF16 以加速,哪些操作(如归一化、softmax)应保留 FP32 以保持数值稳定。

PyTorch 中使用 AMP 的基本流程:

import torch
from torch.cuda.amp import autocast, GradScaler

model = MyModel().cuda()
optimizer = torch.optim.AdamW(model.parameters())
scaler = GradScaler()  # 用于 FP16 的梯度缩放,BF16 不需要

for data in dataloader:
    optimizer.zero_grad()
    with autocast():  # 自动将支持的操作转为 FP16/BF16
        output = model(data)
        loss = criterion(output)
    scaler.scale(loss).backward()  # 缩放loss并反向
    scaler.step(optimizer)         # 更新参数(自动unscale梯度)
    scaler.update()                # 更新scale因子
  • autocast 上下文管理器自动将符合条件的大型矩阵乘法、卷积等操作转为 FP16/BF16 执行,而 LayerNorm、Softmax 等操作保持 FP32。

  • GradScaler 用于 FP16 训练时的梯度缩放,防止小梯度在 FP16 下下溢。如果使用 BF16,则不需要 GradScaler,因为 BF16 的动态范围与 FP32 一致,不会出现下溢问题。

AMP 的好处:无需手动修改模型代码,开箱即用;自动规避易出问题的操作;与 Tensor Cores 无缝配合。

AMP 中的梯度缩放因子如何动态调整?

在 FP16 训练中,梯度值可能小到 FP16 无法表示(约 6×10⁻⁸ 以下),这会导致梯度下溢变为 0,模型无法更新。梯度缩放(Loss Scaling)正是为解决此问题:将 loss 乘以一个较大的缩放因子,反向传播时梯度也被放大,从而被推入 FP16 的可表示范围。更新参数前,再将梯度除以相同的缩放因子还原。

动态缩放策略:缩放因子并非一成不变。如果设得太小,无法完全避免下溢;设得太大,可能导致梯度上溢(变成 Inf/NaN)。PyTorch 的 GradScaler 会自动动态调整缩放因子:

  1. 初始缩放因子设为较大的值(如 2¹⁶ = 65536)。

  2. 每个训练步,scaler.scale(loss).backward() 产生缩放后的梯度。

  3. scaler.step(optimizer) 会检查缩放后的梯度是否包含 Inf/NaN。如果没有溢出,则正常更新参数,并视情况增大缩放因子(比如连续 2000 步无溢出,将缩放因子乘以 2.0)。

  4. 如果检测到溢出(梯度中有 Inf/NaN),则跳过本次更新,并将缩放因子减小(如乘以 0.5)。这步很关键,既保护了模型参数不被破坏,又自适应地找到了安全的缩放范围。

  5. scaler.update() 完成一次动态调整循环。

为什么动态调整重要:训练过程中,梯度分布会变化。训练初期梯度可能较大,需要较小的缩放因子;训练后期梯度变小,需要较大的缩放因子。动态调整让模型始终工作在安全又高效的缩放区间,无需人工干预。

BF16 的情况:BF16 的指数位与 FP32 相同,动态范围极大,几乎不会发生梯度下溢。因此使用 BF16 时,不需要 GradScaler,直接 loss.backward() 即可。这也是 BF16 相比 FP16 的一大优势。