Pytorch-Tricks

写在前面:Pytorch 内容繁杂,实验中用到的技巧与机制若不及时记录,容易遗忘并导致重复造轮子。本文档用于记录 Pytorch 的相关技巧与机制,将持续更新,方便后续查阅。 动态计算图 缘起 在一次实验日志中,发现一个固定的 CNN 模块推理时间存在较大波动。排查后发现是数据流水线卡顿或 GPU 内存管理抖动所致,但也因此注意到了 Pytorch 的动态计算图机制。 Pytorch 在每次前向传播时都会重新构建计算图,这意味着每次前向传播都会产生额外开销,尤其在模型结构复杂或输入数据变化较大时,可能导致推理时间波动。 具体而言: 每次前向传播,即便模型结构未变,Pytorch 仍会依据当前输入数据和模型结构重建计算图。 若输入数据形状或模型结构发生变化,计算图会重新构建,进而引发性能波动。 可能导致性能波动的情况 1. 条件分支结构 模型中包含数据依赖的条件分支(如 if 语句)时,不同前向传播可能执行不同计算路径,导致计算图结构变化。 class BadEncoder(nn.Module): def forward(self, x): if x.sum() > 0: # 数据依赖的条件 x = self.heavy_conv(x) else: x = self.light_conv(x) return x 2. 变长序列 处理变长序列(如 RNN 或 Transformer)时,每次输入长度可能不同,导致计算图结构变化。 class BadRNN(nn.Module): def forward(self, x): for i in range(x.size(1)): # 依赖输入序列长度 x = self.rnn_cell(x[:, i, :], x) return x 3. 其他情况 动态形状 随机性操作 稀疏操作 解决方案 使用静态计算图:若模型结构固定,可借助 TorchScript 或 ONNX 转换为静态计算图,避免每次前向传播都重新构建。 Padding + Mask:对于变长序列,采用填充至统一长度,并利用 mask 标识有效部分,保持计算图结构一致。 模型训练/推理速度分析与优化 背景 修改模型与训练代码后发现速度明显下降。排查发现数据加载耗时过长,遂将 DataLoader 的在线预处理改为离线处理。但模型推理速度仍慢于基线,后续需进一步优化。 ...

April 7, 2026 · 4 min