torch.profiler.profile默认配置无效,因仅录首batch、不捕获cpu阻塞、无warmup导致编译时间干扰;需设schedule(wait=1,warmup=1,active=3)、activities=[cpu,cuda]、record_function精准埋点。

torch.profiler.profile 默认配置几乎无法定位真实瓶颈——它只录第一个 batch,不显示数据加载耗时,也看不出 tensor 拷贝是否在拖慢训练。必须手动配齐关键参数、埋点位置和分析视角,才能让 profiler 输出可行动的结论。
为什么默认 profile 什么也看不出来?
常见现象:GPU 利用率低、训练慢,但 prof.key_averages().table() 里只有 aten::add 或 aten::copy_ 排前列,DataLoader.__next__ 几乎不出现。
- 没开
ProfilerActivity.CPU:profiler 只抓算子,不记录 Python 层阻塞(比如__next__等 I/O) - 没设
schedule:默认只录 step 0,此时 CUDA kernel 还没编译完,把编译时间误当计算瓶颈 -
record_shapes=False:漏掉 batch size 不一致导致的 kernel 重编译,# Calls异常高却查不到根因 -
profile_memory=False:显存峰值卡在某个 layer 输出上,但你只能看到 OOM,看不到哪层突然暴涨
必须设置的 schedule 和 activity 组合
工业级可用的最小配置是:
schedule=torch.profiler.schedule(wait=1, warmup=1, active=3) activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]
这个组合确保:
跳过 step 0(初始化抖动),step 1 触发 kernel 编译(warmup),只分析 step 2–4 的稳定态;同时捕获 CPU 等待、GPU 计算、跨设备拷贝(如 cudaMemcpyAsync)三类关键事件。
- 若用 ROCm 设备(如 MI300X),
ProfilerActivity.CUDA仍有效,但需确认torch.__config__.show()含rocm字样,否则静默降级为 CPU-only -
wait=0或active=1容易把冷启动噪声当瓶颈;active>5会显著拖慢训练,且增加分析噪音
record_function 埋点必须包住实际逻辑,不是模型调用本身
只写 with record_function("forward"): model(inputs) 是无效的——profiler 会把 model.forward 内部所有算子平铺归到 “forward”,无法区分数据预处理、loss 计算、梯度裁剪等环节。
- data_loader 必须包裹在循环体内:
for batch in train_loader:→with record_function("data_load"):→batch = next(train_loader) - forward 段要包含完整前向+loss:
with record_function("forward"):→pred = model(x)→loss = criterion(pred, y) - backward 段必须含
loss.backward()和optimizer.step()全流程,否则梯度更新耗时被拆散到aten::add、aten::mul等底层算子中 - 别在
train_loader外层加record_function,否则__next__耗时会混进 “forward” 统计
看输出时优先盯 Self CPU time total 和 GPU Kernel Utilization 曲线
prof.key_averages().table(sort_by="cpu_time_total") 容易误导:看到 aten::copy_ 排第一就去删 .clone(),其实根因可能是输入 tensor 在 CPU、模型在 GPU,每次 forward 都强制同步拷贝。
- 先看
Self CPU time total(非累计),排除被子调用撑高的父项;若某算子# Calls达 1000+(如index_select),大概率是 for 循环写在了 GPU 上 - 打开 TensorBoard 的
GPU Kernel Utilization曲线,若频繁跌零,说明 GPU 在等 CPU 数据 —— 此时回看data_load区域耗时是否远超forward -
with_stack=True只在查自定义模块或封装层时启用,它会让 profiler 明显变慢,日常分析建议关掉
真正难的是把 trace 里的算子耗时,映射回你写的那几行 Python 逻辑——这要求 record_function 埋点位置精准、命名语义清晰,而不是靠猜。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











