多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

大模型训练性能瓶颈如何定位?用Profile揪出GPU利用率低的真凶

大模型训练性能瓶颈如何定位?用Profile揪出GPU利用率低的真凶 训练一个大模型跑了半天发现 loss 不掉或者 GPU 利用率一直上不去卡在某个数值下不来这种体验我相信做训练的人都不会陌生。很多时候大家的第一反应是“显卡不够好”“显存不够大”但真正用 Profile 扫过一遍之后就会发现瓶颈往往藏在你想不到的地方——数据加载、CPU 预处理、某个算子的实现甚至是一个看似不起眼的参数设置。这篇博文我就用实际跑训练的经验讲讲怎么用 Profile 把大模型训练里的性能瓶颈一步步揪出来以及拿到 profiling 结果之后该怎么读、怎么定位、怎么优化。这篇文章适合正在做大模型训练的算法工程师、平台工程师也适合刚入门分布式训练、想搞清楚“训练为什么这么慢”的读者。我不会只贴命令会把每个环节背后的判断逻辑讲清楚保证你看完能直接在自己的训练任务里用起来。1. 为什么说大模型训练的性能瓶颈得靠 Profile 才能“揪”出来1.1 训练变慢先别急着骂“显卡”大模型训练的链路非常长。从数据读取、样本预处理、tokenizer到数据上 GPU、前向传播、反向传播、梯度同步、优化器更新每一环都可能成为瓶颈。我见过很多次这样的情况GPU 利用率只有 40% 左右但 nvidia-smi 显示显存占用很高于是大家默认“显存不够”其实真正的原因是 CPU 侧的 dataloader 处理速度跟不上GPU 一直在空等数据。单纯靠 nvidia-smi 或者任务管理器那种粗粒度的监控能看到“GPU 在干活”或者“没干活”但看不到“时间到底花在哪了”。哪怕你盯着top、nvtop看一整天也很难定位到一个具体算子或者一段具体代码。这时候就需要 Profile也就是性能画像去把训练过程中每一个环节的耗时、调用关系、资源占用情况完整记录下来。Profile 的本质是“把时间量化到函数和算子级别”。它不是简单告诉你程序跑了多久而是告诉你每一个阶段花了多久、占比多少、有没有在等待、等待的对象是谁。这个信息对定位大模型训练性能瓶颈来说几乎等于一张藏宝图。1.2 普通的计时和性能画像Profile到底差在哪很多人一开始会尝试用最朴素的方式统计时间比如在训练循环里手动记录每个 step 的起止时间或者在 dataloader 的__getitem__里打time.time()。这种计时方式能提供宏观概念比如“一个 step 大概 3 秒”但它有一个严重的问题它只能告诉你某个代码块的总耗时不能告诉你这个耗时是被 CPU 计算占用的还是在等待 GPU 返回或者是卡在 IO 上。Profile 工具会做更深层的插桩。以 PyTorch Profiler 为例它既能记录 CPU 端每个 PyTorch 算子的调用时间也能通过 CUPTI 接口拿到 GPU kernel 的实际执行时间和队列等待时间还能把两者在时间轴上对齐显示 CPU 下发算子的动作和 GPU 实际执行之间的因果关系。这就是为什么 Profile 能精准定位瓶颈——因为它还原了整个训练过程的“时间线”而不仅仅是几个时间点。打个比方普通计时像是你只知道从家到公司通勤花了 1 小时Profile 相当于给这段通勤装了摄像头你能看清是地铁等了 20 分钟、路上堵了 25 分钟、还是电梯排队花了 10 分钟。没有这层细粒度数据你做优化就只能靠猜。1.3 Profile 在大模型训练里的特殊价值大模型训练相比普通深度学习任务有几个更复杂的地方这让 Profile 的价值更加突出。第一模型结构复杂。Transformer 类模型里Attention、LayerNorm、Embedding、FFN 各有各的计算特性和数据访存特性。如果不做 profiling你根本不知道 attention 相关代码是不是真的像理论上那样占据了主要耗时。第二显存压力大。大模型训练普遍用到混合精度、梯度检查点gradient checkpointing、ZeRO 等策略这些策略都会改变显存和计算之间的平衡。显存占用情况、中间激活的大小、碎片化程度单凭感觉是估不准的必须用 memory profiling 配合计算 profiling 一起看。第三分布式通信不可忽视。多卡训练时梯度同步AllReduce、参数广播等通信操作的时间占比会随着卡数上升而上升。通信和计算有没有重叠、通信是不是阻塞了下一个 step这些问题只有通过 profile 去观察时间线上的 NCCL kernel 才能回答。这也是为什么我一直强调大模型训练做性能优化第一步永远是 profiler而不是拍脑袋换配置。2. 工具选型该用哪个 Profile 工具别一上来就 Nsight不少读者一听到 Profile第一反应可能是 Nsight 系列但在实际训练场景里工具的选型取决于你在优化链路的哪个阶段。我的经验是从 PyTorch Profiler 入手再决定要不要上 Nsight。2.1 PyTorch Profiler最直接的入口PyTorch 从 1.8 版本开始集成了torch.profiler这也是我现在最常用的工具。它的优势在于不需要重新编译模型和训练代码天然集成能同时拿到 CPU 算子耗时、GPU kernel 耗时、显存分配记录、shape 信息可以直接导出 trace 文件供 TensorBoard 或 Chrome tracing 查看。一个最基本的用法是这样import torch from torch.profiler import profile, ProfilerActivity, tensorboard_trace_handler def train_step(): # 假设这里是你原本的一个训练 step x torch.randn(512, 2048, devicecuda) w torch.randn(2048, 2048, devicecuda) y torch.matmul(x, w) loss y.sum() loss.backward() with profile( activities[ProfilerActivity.CPU, ProfilerActivity.CUDA], scheduletorch.profiler.schedule(wait1, warmup1, active3, repeat0), on_trace_readytensorboard_trace_handler(./log) ) as prof: for step in range(10): train_step() prof.step()跑完之后./log目录下会生成包含 trace 事件的文件夹用 TensorBoard 加载这个目录即可可视化。这段代码是我建议初学 profile 的人最先跑通的例子——它花不了几分钟但能让你对 profiling 到底输出什么有一个直观认识。2.2 Nsight Systems 和 Nsight Compute 的分工当 PyTorch Profiler 定位到某个 kernel 确实是热点但你想进一步深挖 kernel 内部为什么慢比如访存密集型、计算密集型、还是 launch bound就需要上 NVIDIA 官方的 Nsight 工具。Nsight Systemsnsys是系统级分析工具专门看进程级的 CPU/GPU 活动、CUDA API 调用、Memcpy 操作、NCCL 通信等。它比 PyTorch Profiler 更底层能看到 PyTorch 框架之外的系统行为。Nsight Computencu则是 kernel 级分析工具能给出一个 GPU kernel 的占用率occupancy、寄存器使用、访存带宽、指令混合等详细信息。但我要提醒一句不建议一上来就开ncu因为它对性能的影响很大而且一次只分析一个 kernel需要你先确定要分析哪个 kernel 才有意义。正确的使用路径是先用 PyTorch Profiler 找到热点 kernel再用ncu对准这个 kernel 做深入分析。另外nsys profile和 PyTorch Profiler 是可以配合使用的。nsys主要回答“系统层面卡在哪”PyTorch Profiler 主要回答“模型算子层面卡在哪”。我在实际案例里绝大多数瓶颈用 PyTorch Profiler 就已经能定位了Nsight 用于最后的微观调优。2.3 轻量级辅助手段nvtop、nvidia-smi、自定义计时器这里想给那些连 Profiler 都不方便装的环境提供一套“穷人版”方案。虽然不如专业 profiler 精细但在很多排查场景里已经足够。nvidia-smi配合-l 1每秒刷新能看 GPU 利用率、显存、温度。nvtop是终端里的交互式 GPU 监控类似top。这些工具能告诉你“GPU 有没有在满负荷跑”但无法告诉你“哪个算子在拖后腿”。所以更进一步的轻量做法是自己在训练代码里给关键环节加计时器import time class TimeIt: def __init__(self, name): self.name name def __enter__(self): self.t0 time.time() def __exit__(self, *args): print(f{self.name}: {time.time() - self.t0:.3f}s) with TimeIt(dataloader): batch next(iter(loader)) with TimeIt(forward): loss model(batch)这种计时器的价值在于它能帮你快速切分宏观瓶颈比如“原来 2.5 秒里 1.2 秒在等数据”。知道这一步之后你再用专业 profiler 去定位细节就有的放矢了。我建议每个训练脚本里都保留一套这样轻量的计时工具作为日常 sanity check。工具分析粒度开销上手难度适用场景torch.profiler算子级/GPU kernel中等可控制低日常训练首选的 profilernsys systems系统级/进程级较低中系统调度、通信、IO 问题ncu computekernel 指令级很高高定位热点 kernel 内部瓶颈nvidia-smi / nvtopGPU 整体状态极低极低快速排查 GPU 是否在空等3. 实操给训练脚本动个小手术跑通一次完整的 Profiling3.1 改造训练循环加上 Profiler 上下文很多 training script 的核心循环写得很“结实”直接在原有循环外面包一层 profiler 是可以的但要注意prof.step()的调用位置。step()告诉 profiler 一个训练 step 的边界profiler 根据它划分 iteration才能准确计算单个 step 的平均耗时。实际的改造建议是写一个开关变量只在需要 profiling 的时候才开启平时训练完全不触发import argparse import torch from torch.profiler import profile, ProfilerActivity, tensorboard_trace_handler parser argparse.ArgumentParser() parser.add_argument(--profile, actionstore_true) args parser.parse_args() profiler None if args.profile: profiler profile( activities[ProfilerActivity.CPU, ProfilerActivity.CUDA], scheduletorch.profiler.schedule(wait1, warmup1, active2, repeat1), on_trace_readytensorboard_trace_handler(./prof_log), record_shapesTrue, profile_memoryTrue, ) profiler.start() for step, batch in enumerate(train_loader): # 原有训练代码 loss train_one_step(batch) if profiler: profiler.step() if step 10: break if profiler: profiler.stop()这里有几个细节值得注意。第一break的时机要留够 schedule 需要的步数。我设了wait1, warmup1, active2, repeat1加上warmup额外消耗至少需要跑 5 个 step 以上。第二不要拿第一个 step 的数据当参考因为 CUDA 上下文初始化、cuDNN autotune 都在这个时候发生会严重拉高耗时。wait参数就是用来跳过这一步的。3.2 四个关键参数的取舍决定 Profiling 质量schedule函数有四个参数理解它们比抄代码更重要。wait正式采集前跳过的步数。前面提到过第一个 step 有 CUDA 初始化、cuDNN 搜索等冷启动开销这些不是稳态训练时的情况必须跳过。warmup预热的步数。Profiler 内部的缓存机制、CUPTI 的 hook 需要几次迭代才能进入稳定状态如果不预热前面几步采集的数据会有失真。active真正采集数据的步数。这个值不是越大越好。Profiler 会产生大量 trace 数据active 步数太大会让 trace 文件无比庞大反而难以分析。我一般设置 2~5 步。repeat上述 [wait, warmup, active] 过程重复几次。训练通常是稳态过程采集 1 轮就够了除非你想观察不同 batch 下的波动。另外两个常被忽略的开关是record_shapes和profile_memory。record_shapesTrue会记录每个算子输入输出的 shape 信息这对发现 shape 意外变化导致的重计算非常有帮助。profile_memoryTrue会额外记录显存分配和释放的轨迹但会让开销明显上升建议在需要排查显存问题时才开启active 步数也要相应减少。3.3 看结果TensorBoard 与 trace 文件的基本解读采集完成后./prof_log下会出现形如profiler_log_20240601_120000_worker0.pt.trace.json的文件。用 TensorBoard 查看tensorboard --logdir ./prof_log然后在浏览器里打开http://localhost:6006找到 Profile 标签页。重点看两个视图第一是Overview 视图它会直接给出一张表格列出 kernel 总耗时、CPU 侧耗时、GPU 侧耗时、利用率等关键数字。这个视图适合粗筛能立刻看出是 CPU 还是 GPU 主导了时间。第二是Trace 视图也就是火焰图/时间线。这里能看到每个线程、每个进程在各时间片执行的算子和 kernel颜色区分 CPU 和 GPU 活动。拖动可以精确定位到某个算子看它的 start 时间、duration、所在设备。如果你习惯用 Chrome也可以直接把.pt.trace.json拖进chrome://tracing或者 Perfetto 里查看效果类似。我个人更偏好用 Perfetto因为它在缩放和搜索算子名时更流畅。4. 数据解读从一堆火焰图里精准定位瓶颈拿到 profiler 结果之后最容易犯的错误是被一堆花花绿绿的火焰图带偏到处点点看。我建议按下面这套顺序来读效率最高。4.1 先看全局三件事GPU 利用率、平均步耗时、kernel 时间占比不管 trace 里有多少细节第一步永远是回答三个宏观问题一个 step 平均耗时多少GPU 利用率GPU 上有 kernel 执行的时间占比是多少CPU 端总耗时和 GPU 端总耗时相比谁更大这三个数字能直接把瓶颈分类到大致方向。如果 GPU 利用率低于 60%大概率是 CPU 下发算子太慢或者数据加载卡住了 GPU。这种情况你去优化算子实现是没用的得先解决“GPU 吃不饱”的问题。如果 GPU 利用率看起来很高90% 以上但 step 还是慢那问题就在 GPU kernel 本身的效率或显存带宽上需要继续做算子级分析。如果 CPU 端耗时明显高于 GPU 端说明训练循环被 CPU 操作卡住了常见的元凶包括 dataloader 中过重的预处理、频繁的 CPU-GPU 数据拷贝、Python 层过多的同步点。4.2 算子级分析分清 self time 和 total time在 PyTorch Profiler 的表格视图里每个算子都有两个时间维度Self CPU time / Self CUDA time 和 Total CPU time / Total CUDA time。我见过不少新手直接看 Total 时间排序发现某个算子总耗时最长就认为它是瓶颈。这是不准确的。Total time 包含了它调用的子操作Self time 才是这个算子本身真正执行的时间。如果两个算子 Total 都很大但其中一个 Self 很小说明它只是“组长”把时间花在了调用别人上真正要优化的是 Self 很大的“干活的人”。举个例子如果你发现 Attention 相关算子 Total 耗时很高但展开后 Self 并不大内部主要是bmm、softmax、bmm三段在耗时那你优化的方向就应该是聚焦到这三个底层的 GPU kernel考虑换 FlashAttention 这种融合实现。如果只看 Total可能就稀里糊涂开始调 Attention 的实现结构反而把时间花错地方。4.3 数据管线、显存和通信这三个“隐形瓶颈”怎么看算子表格能定位计算瓶颈但数据加载、显存分配、通信等待这三类问题需要专门去看。数据加载是否成为瓶颈一种有效方法是观察 GPU kernel 时间轴上的间隙。如果 GPU 的 kernel 之间频繁出现大段空白同时 CPU 线程的 dataloader 相关操作处于活动状态那基本可以判定是数据管线跟不上。更直接的做法是分开测一次纯 dataloader 的时间跑一个 epoch完全不调用loss.backward()只看取 batch 的耗时。如果这个耗时已经超过训练步耗时的一半数据侧必然有问题。显存问题的特征是显存占用不断增长、或者出现显存不足崩溃。通过profile_memoryTrue采集的数据可以在 TensorBoard 里看每个 tensor 的分配时间点和生命周期。还可以结合torch.cuda.memory_summary()查看内存池状态确认是否有大量碎片化。大模型训练里显存碎片化是一个很容易被忽视的问题激活值的反复分配释放会让缓存池产生很多小空洞导致即使理论显存够用实际也会 OOM。通信瓶颈集中出现在多卡训练中。在 trace 时间线里搜索NCCL关键字如果看到 AllReduce 的 kernel 占了大段时间且前后计算 kernel 有等待关系没有重叠就说明通信没有和计算重叠。这时候要考虑使用更高效的通信后端、调整梯度累积策略或者用 ZeRO 系列优化器减少通信量。5. 实战案例7B 模型训练从 65% GPU 利用率到 91%这里我拿一个虚拟但非常典型的案例来完整复盘一遍 profiling 驱动优化的流程。假设我们在单机 8 卡上训练一个 7B 参数量的模型序列长度 2048每卡 batch size 1用的是混合精度训练。5.1 第一轮 Profile瓶颈原来在 CPU 数据侧训练一开始我们观察到 GPU 利用率只有 65% 左右step 耗时约 5.2 秒。按照前面的分析流程先跑一轮 PyTorch Profiler。结果 Overview 视图显示GPU kernel 总耗时只有 3.2 秒但一个 step 的墙钟时间是 5.2 秒。trace 时间线上 GPU 有大量空闲区间CPU 端的 dataloader 相关操作collate、to(device)、数据处理函数Occupancy 很高。进一步看数据加载统计发现单次__getitem__平均耗时 120ms而每 step 只取 1 个样本。问题显而易见数据预处理太重了。代码里有一个实时做数据增强的逻辑对每张训练样本做了多次随机裁剪和高斯模糊这些操作全在 CPU 上同步执行。优化措施把num_workers从 4 调到 16prefetch_factor从 2 调到 8开启pin_memoryTrue把 dataloader 的 batch 放进页锁定内存减少 CPU 到 GPU 拷贝时间把部分数据增强逻辑改为 GPU 上执行用torchvision.transforms的 GPU 版本或者干脆用离线预增强把增强后的数据缓存到内存/磁盘。改造后再跑 profileGPU 利用率升到 88%step 耗时降到 3.6 秒。这时候 GPU kernel 总耗时接近 3.1 秒说明数据侧已经不是主要瓶颈了。5.2 第二轮 Profile算子层面才是真正的“大户”GPU 利用率上来以后优化重点转向 kernel 效率。再跑一轮 profile把算子按 Self CUDA time 排序。发现 Attention 里的bmmtorch.bmm两个矩阵乘法和softmax加起来占了 GPU 总耗时的 41%明显异常偏高。原因是我们用的是 PyTorch 原生的 attention 实现包含 mask 的torch.masked_fill和softmax中间会产生大量不必要的内存读写。另外record_shapes还暴露了一个问题注意力分数矩阵的 shape 竟然是 (8, 32, 2048, 2048)也就是说 mask 是在 batch 内动态广播的一方面占显存另一方面也让访存变得碎片化。优化措施把原生 attention 换成 FlashAttention通过torch.nn.functional.scaled_dot_product_attention它在 kernel 内部完成 attention 计算和 softmax不把中间矩阵写回显存mask 提前预计算成常量并在每个 GPU 上缓存避免每步重复生成确保模型以 bf16 精度运行减少显存带宽压力。效果立竿见影attention 相关 kernel 耗时从 GPU 总耗时的 41% 降到 18%step 耗时从 3.6 秒降到 2.3 秒。GPU 利用率也顺势上升到 92% 左右。5.3 第三轮 Profile多卡场景下的通信与显存平衡继续 profiling 时我们把视角切到多卡通信。在 trace 时间线里搜索NCCL发现每次反向传播结束后会立即执行 AllReduce而且这段时间 GPU 计算基本是空的没有和通信重叠。8 卡场景下通信耗时约占整个 step 的 15%。同时显存 profile 显示中间激活值峰值达到每卡 14GB已经接近我们 A100 80G 卡的 17% 左右。显存还有余量但需要为梯度检查点等策略留出空间。优化措施开启梯度累积增大等效 batch size减少通信频率例如每 4 个微批次再进行一次梯度同步通信次数减少为原来的 1/4把训练模式改成 gradient checkpointing虽然增加了约 10% 的计算开销但显存峰值下降了约 30%给更大的 micro batch size 留出了空间提速网络确保多卡之间用 NVLink 通信避免走 PCIe。调整后通信时间占比降到 4%step 耗时进一步降到 1.9 秒。5.4 三轮优化后的效果对比指标初始状态第一轮后第二轮后第三轮后GPU 利用率65%88%92%94%单 step 耗时5.2s3.6s2.3s1.9sattention 相关耗时占比约 35%约 33%18%17%通信耗时占比约 8%约 8%约 10%4%显存峰值约 10GB约 10GB约 9GB约 6GB一个非常直观的结论是同样的硬件不变只是靠 profiling 驱动的优化整个训练吞吐提升了接近 2.7 倍。这轮优化的每一步决策都来自 profiler 给出的量化解剖而不是拍脑门。6. 常见问题与避坑实录Profile 不是万能的但比瞎猜强一百倍6.1 Profiler 本身的开销会影响结果吗会但没有想象中那么大关键看你怎么配置。Profiler 的插桩本身会占用一定的 CPU 和 GPU 时间CUPTI 的回调机制在部分 GPU 上还会影响 kernel 执行效率。因此采集到的绝对耗时和不开 profiler 时的实际训练耗时存在偏差。但这不是大问题因为 profiling 的目的是找到瓶颈的“占比”和“相对关系”而不是追求绝对精准的微秒级数据。我常用的策略是减小active步数到 2~3 步减少 trace 数据量平时不 profiling 训练只有需要优化时才开一小段。如果你发现开 profiler 后显存占用异常升高记得检查profile_memory是否开启这个开关会让显存记录额外分配缓存建议只在排查显存问题时开启。6.2 分布式训练要不要每张卡都开 Profile很多人会想当然地在所有 rank 上开启 profiler结果得到一个巨大无比、几乎无法打开的 trace 文件。实际上分布式训练里各卡的训练路径是一致的没必要全开。我推荐的方案是只在一个 rank 上开 profiler通常是 rank 0用来分析计算和数据加载如果要看通信耗时则额外在 rank 0 和 rank 1 上各开一次对比两个 rank 之间的 NCCL 同步情况。注意分布式 profiler 必须保证所有 rank 的 schedule 一致否则会出现不同步的采集窗口导致某些 rank 等待它自己的 profiler 计时器反而干扰了通信测量。还有一点容易踩坑开启 profiler 的 rank 会在 profiler 停止时花时间写 trace 文件其他 rank 可能已经进入下一个训练循环造成短暂的不同步。所以 profiling 结束后最好直接跑够预定步数退出不要在不该退出的位置让部分 rank 干等。6.3 几个容易忽略的细节warmup、小规模跑、基线的价值最后分享三个我在实践中总结的小建议。第一一定要跑 warmup 再分析。甚至可以说warmup 比 active 本身更重要。如果直接对第一个 step profiling你会看到一堆 cuDNN autotune 和 CUDA 初始化操作占据大量时间然后误判瓶颈在卷积或者某个库加载上白白浪费时间。第二小规模跑也能发现问题。很多人担心 7B 模型没办法在本地快速验证 profiling 流程其实完全可以先用一个小模型比如几十 MB 的 GPT-2 规模在小数据集上跑通 profiling 的完整流程确认采集、导出、分析链路正常再上大模型。链路本身是一样的小规模环境里踩过的坑上了大规模基本能避开。第三建立 profiling 基线。我建议在项目初始化时、每次改动模型结构或训练配置后都保存一份 profiling 摘要。后面再出性能问题时和基线对比就能快速定位是哪一次改动引入了瓶颈。这个习惯在我维护长期训练任务时帮了很大忙。你可以把每次 profile 的关键指标step 耗时、GPU 利用率、Top 5 耗时算子及其占比存成一个简单的 JSON 或 Markdown 表格放进项目仓库里随代码一起管理。根据我自己的经验用 profiler 定位性能问题最难的往往不是工具本身而是敢不敢对“我以为的瓶颈”下手。很多时候打开 trace 之后你会发现实际情况和直觉完全相反。所以我的建议是遇到训练变慢先跑一轮 profile 再说话。这个习惯几乎能帮你省掉一半的无效加班。
返回列表