第一次真正跑大模型训练的时候,我盯着机房监控界面。风扇声音很大,GPU 利用率却在那儿晃悠。心里真的会“咯噔”一下——钱在烧,速度却没上去。

       做训练久了你会发现,问题不复杂。显存顶不住,训练就卡住;GPU 吃不饱,钱就白花。很多人一上来就加卡、堆算力。要我说,这种方式最贵,也最懒。

真正该动手的,是显存、并行方式,还有精度。


一、显存:训练能不能跑起来,先看它

显存为什么老是不够?

Transformer 结构里,显存被几样东西吃掉:

  • 参数

  • 梯度

  • 优化器状态(Adam 里那两个动量)

  • 中间激活值

很多人只盯参数量。其实优化器状态才是大头。
一个 10G 参数模型,训练时显存至少翻四倍。参数一份,梯度一份,动量两份。显存条一下子就红了。

梯度检查点:用算力换空间

        第一次开 Gradient Checkpointing,我心里是犹豫的。担心变慢。结果一跑,显存直接下来三四成。

它的逻辑很简单——中间结果不存,反向传播时重算一遍。

from torch.utils.checkpoint import checkpoint

def forward(self, x):
    x = checkpoint(self.transformer_block, x)
    return x

算力会多花一点。大概 10% 左右。
但如果你显存已经 OOM,那点时间根本不算事。

8-bit 优化器:显存肉眼可见地降

Adam 很稳,但它太占地方。后来我换成 8-bit 版本,监控面板上的显存条往下一缩,挺解压。

import bitsandbytes as bnb
optimizer = bnb.optim.Adam8bit(model.parameters(), lr=1e-4)

精度几乎没掉。显存少一半。很值。

ZeRO:真正的大杀器

多卡训练时,如果每张卡都存完整模型,那是浪费。
Microsoft 在 DeepSpeed 里做的 ZeRO,本质就是“拆”。

优化器拆开。
梯度拆开。
参数也拆开。

       卡多的时候,ZeRO-3 的效果很明显。原来 8 卡都顶不住的模型,突然就能跑了。那一刻真的挺爽。


二、并行:别让 GPU 发呆

有一次我看 nvidia-smi。GPU 利用率 58%。风扇转得飞快。机器在响。卡却在等数据。

这种感觉很难受。

数据并行

最常见。每张卡一份模型。喂不同数据。
简单。好用。
模型一大就爆显存。

模型并行

把模型切开。Attention 在一张卡。MLP 在另一张。
能跑更大的模型。
通信复杂。调一次参数,能折腾一下午。

张量并行 + Pipeline

NVIDIA 在 Megatron 里搞的那套思路,很适合百亿级模型。

矩阵拆分。
层级分段。
算力能铺得更均匀。

不过别盲目上。卡少的时候,复杂并行反而拖慢。


三、混合精度:几乎是白给的优化

第一次切到 BF16,我有点紧张。怕 loss 飘。
结果训练曲线很稳。速度却明显快了。

FP32 太重。
FP16 快,但有时会溢出。
BF16 现在是首选。只要卡支持。

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

with autocast():
    loss = model(input)

scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

显存几乎砍半。
速度上去一截。
这种优化,说实话,属于“不开白不开”。


四、成本到底花在哪

很多人算成本,只算 GPU 单价。
其实真正贵的,是反复试错。

参数不收敛。
数据有问题。
学习率不对。

每一次重新跑,都是几千上万块。

GPU 利用率低于 70%,我都会不舒服。那意味着资源在空转。

要我说,优化前先盯三个东西:

  • 显存占用结构

  • GPU 利用率

  • 通信时间比例

有时候不是模型太大,是数据加载太慢。磁盘在拖后腿。


五、一个真实案例

我们当时训一个 13B 模型。
8 张 24G 卡。
频繁 OOM。利用率只有 60%。

调的过程挺折腾。

开 BF16。
上 ZeRO-2。
加 Gradient Checkpoint。
Batch 从 2 拉到 8,用梯度累积顶上。

几轮之后,GPU 利用率冲到 90% 以上。
显存稳定。
整体时间缩短三成左右。

不是某一个技巧救了场。是组合拳。


六、一些更“值钱”的经验

别一上来就大模型。
小模型先把学习率、数据质量跑顺。

数据干净,比多几亿参数重要。
脏数据会让 loss 乱跳。看着曲线心里发凉。

也别迷信超大 batch。
泛化可能变差。最后还得回头调。

多机训练时,网络拓扑很关键。
NVLink 和 PCIe 差别明显。
跨节点通信慢的时候,你能听见 GPU 在“等”。


大模型训练不是玄学。
也不是堆钱游戏。

显存压一压。
并行铺一铺。
精度降一档。

很多时候,成本能省下一大块。

       要我说,真正厉害的优化,不是参数调得多花哨,而是你能盯着监控面板,知道每一块显存、每一秒时间在干什么。

机器在转。
曲线在降。
那种踏实感,比什么都值钱。

更多推荐