pytorch profiler可实时定位grok推理瓶颈:需pytorch≥1.12、cuda驱动就绪,双验证后安装tensorboard及插件,创建日志目录,在model.generate()前后注入profiler并启用record_shapes、with_stack、profile_memory,禁用autocast,导出trace.json,通过tensorboard的trace viewer查看gpu耗时、层级卡点与显存泄漏。
☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 多模态理解力帮你轻松跨越从0到1的创作门槛☜☜☜

你需要在Grok模型推理服务上线后,立刻看清每个请求在GPU上花了多少时间、卡在哪一层、显存是否暴涨——而不是等用户投诉后再翻日志查半天。
准备TensorBoard Profiler运行环境
确保PyTorch版本 ≥ 1.12且CUDA驱动已就绪,执行nvcc --version和python -c "import torch; print(torch.__version__, torch.cuda.is_available())"双验证。
安装TensorBoard及配套插件:pip install tensorboard torch-utils-tensorboard。注意:不要用torch.utils.tensorboard旧包名,新版本已统一为torch.utils.tensorboard,拼错会导致ImportError: cannot import name 'SummaryWriter'。
创建日志目录:mkdir -p ./grok_profile_logs,该路径后续必须与代码中log_dir参数严格一致。
在Grok推理函数中注入Profiler探针
打开你的Grok服务主文件(如inference.py或app.py),定位到实际执行model.generate()或model.forward()的函数入口。
在函数开头插入Profiler初始化代码:
from torch.profiler import profile, record_function, ProfilerActivityprof = profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA], record_shapes=True, profile_memory=True, with_stack=True, with_flops=True)
在模型调用前后包裹prof.start()和prof.stop(),并在prof.export_chrome_trace("./grok_profile_logs/trace.json")导出结果。这一步漏掉record_shapes=True将无法看到张量维度变化,而with_stack=True缺失则无法定位到具体Python行号。
【必须关闭autocast】若代码中启用了torch.cuda.amp.autocast(),需在profiler作用域内临时禁用,否则trace中CUDA kernel名称会显示为__amp_after_resume等不可读标识,丧失定位价值。
启动TensorBoard并加载追踪数据
终端执行:tensorboard --logdir=./grok_profile_logs --bind_all --port=6006。使用--bind_all确保远程服务器可访问,不加此参数时本地浏览器打不开http://服务器IP:6006。
浏览器打开http://localhost:6006/#pytorch_profiler,首次加载可能需等待10–30秒——这是TensorBoard解析trace.json生成交互式timeline的正常过程,不是卡死。
点击顶部工具栏的「Trace Viewer」标签页,你会看到横向时间轴上密布的彩色条形块,每个块代表一个CUDA kernel或CPU算子;纵向按线程/流分层排列,GPU活动集中在底部几行。
按住鼠标左键拖拽放大某一段高延迟区域,再右键→「View Call Stack」,即可看到该kernel对应到Grok源码中的具体位置,例如transformer_block.py第47行 attn_mask.expand()。
识别三类典型Grok推理瓶颈
方法一:看GPU利用率曲线
在「Overview」面板中观察「GPU Utilization」折线图,若峰值长期低于30%,说明计算未饱和——此时检查batch_size是否过小,或kv_cache未复用导致重复计算。
方法二:抓内存尖峰
切换到「Memory」视图,重点查看「GPU Memory Usage」曲线。若每次推理后显存未回落至基线,而是阶梯式上涨,【说明存在tensor缓存泄漏】,常见于未调用del outputs或torch.cuda.empty_cache()。
方法三:定位长尾算子
在「Operator Stats」表格中按「Self CPU time (ms)」倒序排列,找到耗时TOP5的算子。若出现大量aten::copy_或aten::narrow,说明输入序列padding过度;若aten::bmm占比超40%,则应检查attention head数与GPU SM数量是否匹配。











