1. 从一次训练任务说起:为什么全链路优化比单点提速更值得做
去年年底,我接手了一个具身智能方向的训练任务,基座模型是 GR00T N1.6,硬件是单机八卡 A100 80G 的配置。当时团队的目标很朴素:把训练周期压下来,让迭代速度跟上算法侧的节奏。但真正跑起来之后,问题一个接一个冒出来——单步耗时波动大、GPU 利用率上不去、通信阶段大量空转、显存峰值卡在临界点导致 batch size 上不去。我们试过调 dataloader 的 worker 数、试过换优化器、试过手动调 NCCL 参数,收益都很有限,单步时间最多降个百分之十几,离“训练周期减半”差得远。
后来我们把视角从“单点优化”切换到“全链路优化”,用 LoongForge 这套训练框架对 GR00T N1.6 的训练流程做了一次系统性重构,最终把吞吐拉到了原来的 2.3 倍,训练周期基本砍半。这篇文章就把这次优化的完整思路、关键决策、实操细节和踩过的坑都摊开讲一遍。如果你也在做具身智能模型训练,或者手上有类似规模的多模态训练任务,正在被吞吐和显存卡脖子,那这篇内容应该能给你省不少试错时间。
先说清楚这次优化的核心关键词:LoongForge是训练框架层,GR00T N1.6是被优化的模型,CUDA Graph和通信-计算重叠是两个最关键的底层手段。整篇文章会围绕这四个点展开,但不会只讲概念,而是把每一步为什么这么做、参数怎么算、代码怎么改都讲透。
2. 优化前的基线:先搞清楚瓶颈到底在哪
2.1 基线配置与实测数据
优化之前,我们的训练配置大致是这样的:GR00T N1.6 的视觉编码器用 ViT-L,动作头是 diffusion policy 结构,batch size 设为 64,序列长度 512,混合精度用 bf16,优化器是 AdamW,学习率走 cosine schedule。单机八卡,卡间用 NVLink,节点内通信走 NCCL。
跑出来的基线数据很典型:单步耗时约 1.85 秒,其中前向+反向计算占 1.1 秒,数据加载占 0.15 秒,优化器更新占 0.2 秒,剩下的 0.4 秒基本耗在通信和 kernel launch 的间隙上。GPU 利用率用 nvtop 看,峰值能到 85%,但平均只有 62% 左右,说明有大量时间 GPU 在等。显存峰值 76G,离 80G 的上限很近,导致 batch size 不敢往上加。
提示:做优化之前一定要先建立可靠的基线。我见过太多人一上来就改代码,改完发现不知道是变快了还是变慢了,因为没有对照。基线要记录单步耗时、各阶段占比、GPU 利用率、显存峰值这四个核心指标。
2.2 瓶颈定位:三个被忽视的“隐形开销”
很多人看训练慢,第一反应是“计算不够快”,于是去换更快的卡或者优化算子。但我们用 profiler 抓了一遍 timeline 之后发现,真正的大头不在计算本身,而在三个隐形开销上。
第一个是kernel launch 开销。GR00T N1.6 的结构里有大量小算子,尤其是 diffusion 动作头部分,每步要启动几百个小 kernel,每个 kernel 的启动开销虽然只有几微秒,但累积起来单步就是几十毫秒。第二个是通信-计算串行。原来的流程是:前向算完 → 等所有卡同步 → 反向 → 再同步 → 优化器更新。每次同步都是一次全局 barrier,GPU 在这期间完全空转。第三个是显存碎片。PyTorch 默认的显存分配器在长时间训练后会产生碎片,导致明明有空间却分配不出连续的大块,batch size 上不去。
这三个问题单独看都不致命,但叠加在一起,就把 GPU 利用率从理论上的 90% 拉到了 62%。所以优化的思路很明确:不是去换更快的硬件,而是把这部分被浪费掉的时间抢回来。
3. LoongForge 全链路优化的整体设计思路
3.1 为什么选择全链路而不是单点优化
单点优化的天花板很低。比如你只优化 dataloader,最多把数据加载的 0.15 秒降到 0.05 秒,单步从 1.85 降到 1.75,提升不到 6%。你只优化通信,把 0.4 秒的间隙压到 0.2 秒,提升也就 11%。但如果把计算、通信、显存、调度这几条链路一起优化,让它们互相重叠、互相配合,收益是乘性的,不是加性的。
LoongForge 的设计哲学就是“全链路”。它不是一个单纯的分布式训练库,而是一套覆盖数据加载、计算图捕获、通信调度、显存管理的完整框架。用它来优化 GR00T N1.6,相当于给整个训练流程做了一次系统性的“管道重排”,让原本串行的环节尽量并行,让原本空转的 GPU 尽量满载。
3.2 三条优化主线:CUDA Graph、通信-计算重叠、显存池化
具体到实现上,我们走了三条主线。
第一条是CUDA Graph 捕获。把 GR00T N1.6 的前向和反向计算图整体捕获成一个 CUDA Graph,这样原本几百次 kernel launch 就变成了一次 graph launch,kernel launch 开销直接从几十毫秒降到几毫秒。这条线解决的是“CPU 调度跟不上 GPU”的问题。
第二条是通信-计算重叠。把梯度同步的通信操作和前向、反向的计算操作在时间轴上错开,让通信在后台进行,计算在前台继续。这条线解决的是“GPU 等通信”的问题。实现上用的是 LoongForge 提供的 overlap scheduler,配合 NCCL 的异步通信接口。
第三条是显存池化与碎片整理。LoongForge 内置了一个显存池管理器,把常用的张量形状预分配好,训练过程中复用,避免频繁的 malloc/free 产生碎片。这条线解决的是“显存够但分配不出”的问题,直接让 batch size 从 64 提到了 96。
这三条线不是孤立的。CUDA Graph 捕获之后,通信操作的位置就固定了,这反而让通信-计算重叠的调度更容易做;显存池化之后,张量地址稳定,CUDA Graph 的捕获也更可靠。所以全链路优化的关键,是让这几条线互相配合,而不是各做各的。
4. CUDA Graph 在 GR00T N1.6 上的落地细节
4.1 哪些部分适合捕获,哪些不适合
CUDA Graph 不是万能的,它要求被捕获的计算图在每次迭代中结构完全一致,不能有动态控制流,不能有依赖 CPU 数据的条件分支。GR00T N1.6 的结构里,视觉编码器和动作头的主体部分都是静态图,适合捕获;但数据加载、随机增强、以及一些依赖 batch 内统计量的操作,就不适合放进 graph。
我们的做法是“分段捕获”:把前向+反向的计算部分单独捕获成一个 graph,优化器更新单独捕获成另一个 graph,数据加载和增强留在 graph 外面用常规方式跑。这样既拿到了 kernel launch 优化的收益,又避免了动态部分破坏 graph 的稳定性。
import torch from loongforge.graph import GraphCapture # 前向+反向的捕获 model = GR00TN16(...).cuda() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) capture = GraphCapture(model, optimizer) capture.warmup(dummy_input, steps=3) # 预热,让显存分配器稳定 graph = capture.capture(dummy_input) # 正式捕获 # 训练循环里直接 replay for batch in dataloader: loss = graph.replay(batch) loss.backward() optimizer.step()注意:捕获之前一定要做足够的 warmup。PyTorch 的显存分配器在第一次运行时会有 lazy 分配,如果直接捕获,graph 里记录的地址可能是临时的,replay 时会出错。我们一般 warmup 3 到 5 步,确保显存分配稳定后再捕获。
4.2 捕获过程中的显存与地址稳定性问题
CUDA Graph 最坑的地方是地址稳定性。graph 捕获时会记录所有张量的显存地址,replay 时这些地址必须仍然有效。如果训练过程中有新的张量分配,或者旧的张量被释放,地址就可能冲突。
我们踩过的坑是:优化器更新时 AdamW 会创建动量张量,这些张量如果在捕获之后才分配,就会和 graph 里的地址冲突。解决办法是在捕获之前就把优化器的状态初始化好,让动量张量提前分配。LoongForge 的 GraphCapture 提供了preallocate_optimizer_state参数,打开之后会自动处理这件事。
另一个坑是 batch size 变化。如果训练中途改了 batch size,graph 就失效了,必须重新捕获。所以我们的策略是固定 batch size,用梯度累积来模拟更大的 batch,而不是动态改 batch size。
4.3 实测收益:kernel launch 开销从 40ms 降到 4ms
捕获之后,单步的 kernel launch 开销从原来的约 40 毫秒降到了 4 毫秒左右,降了 90%。这部分收益直接体现在单步耗时上,从 1.85 秒降到了 1.5 秒左右。更重要的是,GPU 利用率的波动变小了,原来因为 kernel launch 间隙导致的利用率掉坑基本消失,平均利用率从 62% 提到了 74%。
5. 通信-计算重叠:让 GPU 不再等 NCCL
5.1 梯度同步的串行代价
分布式训练里,梯度同步是绕不开的。八卡训练时,每次反向传播结束都要做一次 all-reduce,把八张卡的梯度平均。这个 all-reduce 的数据量是模型参数量乘以 2(bf16),GR00T N1.6 大概有 3B 参数,一次 all-reduce 要传 6GB 的数据。在 NVLink 上,理论带宽是 600GB/s,实际能跑到 400GB/s 左右,所以一次 all-reduce 大概要 15 毫秒。听起来不多,但每步都要做,而且做的时候 GPU 在等,累积起来就很可观。
原来的流程是:反向算完 → all-reduce → 优化器更新。all-reduce 期间 GPU 完全空转。我们的优化目标就是把这个 all-reduce 藏到计算后面,让它和反向传播的计算重叠起来。
5.2 LoongForge 的 overlap scheduler 怎么用
LoongForge 提供了一个 overlap scheduler,核心思路是梯度分桶 + 异步通信。把模型的梯度按层分成若干桶,每算完一层的梯度就立刻发起这一桶的 all-reduce,而不是等所有梯度都算完再一起做。这样通信就和后续层的反向计算重叠起来了。
from loongforge.distributed import OverlapScheduler scheduler = OverlapScheduler( model=model, bucket_size_mb=64, # 每个桶 64MB overlap_ratio=0.8, # 目标重叠比例 backend='nccl', ) # 训练循环 for batch in dataloader: with scheduler.overlap(): loss = model(batch) loss.backward() optimizer.step()bucket_size_mb这个参数很关键。桶太小,通信次数多,启动开销大;桶太大,重叠窗口短,重叠效果差。我们实测下来 64MB 是个比较平衡的值,八卡配置下能把重叠比例做到 75% 到 80%。
5.3 重叠比例的计算与调参经验
重叠比例的计算方式是:被通信覆盖的计算时间 / 总通信时间。假设反向传播总耗时 800 毫秒,all-reduce 总耗时 120 毫秒,如果重叠做得好,120 毫秒里有 96 毫秒是藏在计算后面的,重叠比例就是 80%。
调这个参数的经验是:先看反向传播的耗时分布。如果反向传播是均匀的,重叠容易做;如果某一层特别重,通信就会卡在那一层后面。GR00T N1.6 的视觉编码器部分比较均匀,动作头部分有波动,所以我们把桶的划分做得更细,让动作头部分的通信也能找到重叠窗口。
提示:重叠不是越多越好。如果 overlap_ratio 设得太高,通信会抢占计算资源,反而拖慢计算。我们一般从 0.6 开始试,逐步往上加,找到吞吐的拐点。
6. 显存池化与 batch size 提升的实操
6.1 显存碎片的成因与观测方法
显存碎片是个很隐蔽的问题。PyTorch 的默认分配器是 caching allocator,它会缓存已经释放的显存块,下次分配时优先复用。但如果张量形状变化频繁,缓存块的大小对不上,就会产生碎片。表现出来就是:torch.cuda.memory_allocated()显示只用了 60G,但torch.cuda.memory_reserved()显示预留了 78G,再想分配大块就失败了。
观测碎片的方法是用torch.cuda.memory_summary(),看 reserved 和 allocated 的差值。如果差值超过 10G,说明碎片比较严重。我们优化前这个差值是 14G,优化后降到了 3G 以内。
6.2 LoongForge 显存池的配置与复用策略
LoongForge 的显存池管理器思路很简单:把训练中常用的张量形状统计出来,预分配一批固定大小的块,训练过程中所有张量都从这些块里切,不再走系统的 malloc/free。这样地址稳定,碎片自然就没了。
from loongforge.memory import MemoryPool pool = MemoryPool( total_size_gb=78, # 预留 78G 给显存池 block_sizes=[1, 2, 4, 8, 16, 32], # 块大小梯度,单位 GB enable_defrag=True, # 开启碎片整理 ) with pool.scope(): for batch in dataloader: loss = model(batch) loss.backward() optimizer.step()block_sizes的设计有讲究。块大小要覆盖训练中常见的张量尺寸,太小会导致大张量切不出来,太大又浪费。我们统计了 GR00T N1.6 训练中所有张量的尺寸分布,发现大部分集中在 2G 到 16G 之间,所以块大小梯度设成 1、2、4、8、16、32 这六档。
6.3 batch size 从 64 提到 96 的完整过程
显存池化之后,显存峰值从 76G 降到了 68G,多出来的 8G 让我们可以把 batch size 从 64 提到 96。但提 batch size 不是改个数字就完事,还要同步调学习率。我们的做法是线性缩放:batch size 从 64 到 96,是 1.5 倍,学习率从 1e-4 提到 1.5e-4,同时把 warmup 步数从 500 提到 750,避免初期震荡。
提完之后单步耗时从 1.5 秒涨到了 1.9 秒,但处理的样本数从 64 涨到了 96,折算下来每样本耗时从 23.4 毫秒降到了 19.8 毫秒,吞吐提升了 18%。这就是显存优化的间接收益——它本身不直接提速,但通过允许更大的 batch,让整体吞吐上去了。
7. 常见问题与排查技巧实录
7.1 CUDA Graph 捕获失败的典型原因
CUDA Graph 捕获失败是最常见的问题,报错信息往往很模糊。我们整理了几种典型情况和对应的排查方法。
| 报错信息 | 根本原因 | 解决方法 |
|---|---|---|
operation not permitted during capture | 捕获期间有 CPU 同步操作 | 检查代码里是否有.item()、.cpu()、print等操作 |
invalid argument | 张量地址在捕获后失效 | 确保所有张量在捕获前已分配,开启preallocate_optimizer_state |
out of memory during capture | 捕获时显存不足 | 降低 batch size 或开启显存池 |
graph replay mismatch | 输入形状和捕获时不一致 | 固定输入形状,用 padding 处理变长输入 |
排查的时候,先用CUDA_LAUNCH_BLOCKING=1环境变量跑一遍,让报错定位到具体行。然后检查那一行是否有动态控制流或者 CPU 同步。
7.2 通信重叠后 loss 不收敛的排查
通信-计算重叠做不好,最直接的后果就是 loss 不收敛。我们遇到过一次,重叠开启后 loss 在前 100 步正常下降,之后突然震荡。排查下来发现是梯度桶的划分有问题:某一层的梯度被分到了两个桶里,all-reduce 的时候顺序错乱,导致梯度更新不一致。
解决办法是确保每个桶里的梯度来自连续的层,不要跨层切分。LoongForge 的 OverlapScheduler 默认是按层顺序分桶的,但如果手动指定了 bucket 边界,就要注意这个问题。另外,重叠开启后建议把梯度裁剪的阈值调小一点,因为异步通信可能引入微小的数值误差。
7.3 显存池化后的性能回退问题
显存池化不是没有代价的。我们第一次开启后,发现单步耗时反而涨了 5%。排查下来是块大小梯度和实际张量尺寸不匹配,导致大量张量要走 fallback 路径,从系统分配器拿内存,反而更慢。
调整方法是用 LoongForge 提供的analyze_tensor_shapes工具,先跑 100 步,统计所有张量的尺寸分布,然后根据分布重新设计block_sizes。我们调整之后,fallback 比例从 30% 降到了 5% 以下,性能回退消失,还额外拿到了 3% 的提升。
提示:显存池化一定要先分析再配置,不要拍脑袋设块大小。分析工具跑 100 步大概花 2 分钟,但能省下几个小时的调参时间。
8. 优化效果复盘与后续可扩展的方向
8.1 最终数据:吞吐 2.3 倍,周期减半
把所有优化叠加起来之后,最终的数据是这样的:单步耗时从 1.85 秒降到 1.9 秒(batch size 从 64 提到 96),每样本耗时从 28.9 毫秒降到 19.8 毫秒,吞吐提升到 2.3 倍。训练周期从原来的 14 天降到了 6.5 天,基本砍半。
GPU 平均利用率从 62% 提到了 88%,显存峰值从 76G 降到 68G,碎片从 14G 降到 3G 以内。这些数字背后,是 CUDA Graph、通信重叠、显存池化三条线共同作用的结果,单独任何一条都做不到这个效果。
8.2 还能继续压榨的空间
优化到这一步,还有没有空间?有,但边际收益在递减。我们目前看到两个方向:一是把数据加载也纳入 CUDA Graph 的覆盖范围,用 GPU 侧的预处理替代 CPU 侧的 dataloader,能再省 30 到 50 毫秒;二是把优化器更新也做成 graph,目前优化器更新还是单独跑的,如果和反向传播合并捕获,能再省一点 kernel launch 开销。
另一个方向是通信的进一步优化。目前用的是 ring all-reduce,如果换成 tree 或者 hierarchical 的拓扑,在八卡配置下可能还能再压 10% 到 15% 的通信时间。但这个要看具体的网络拓扑,不是所有机器都适用。
8.3 给同类训练任务的迁移建议
如果你手上的任务和 GR00T N1.6 类似,是多模态、大参数量、单机多卡的配置,那这套优化思路基本可以直接迁移。顺序上建议先做显存池化,因为它风险最低、收益最直接;再做 CUDA Graph,收益大但坑也多;最后做通信重叠,因为它对模型结构有一定要求,需要调参。
有一点要提醒:不要一次性把所有优化都打开。我们当时是分三周逐步上的,每周只开一个优化,观察一周的稳定性再上下一个。这样出问题的时候容易定位,也不会因为多个优化互相干扰而找不到根因。训练框架的优化,稳比快重要,尤其是在生产环境里。
最后分享一个我们踩过的坑:CUDA Graph 捕获之后,如果你用torch.save保存 checkpoint,保存的是 graph 外部的模型状态,不是 graph 内部的。恢复的时候要先加载模型状态,再重新捕获 graph,不能直接保存 graph 对象。这个坑我们卡了大半天才搞明白,希望你能绕过去。