ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

深入理解model.eval()与torch.no_grad():推理阶段显存与速度优化实战

深入理解model.eval()与torch.no_grad():推理阶段显存与速度优化实战 1. 推理阶段为什么必须区分 train 与 eval从一次显存爆掉说起很多人第一次把 PyTorch 模型搬到线上时都会遇到一个诡异现象训练时 batch size 开到 64 都稳推理时 batch size 只开到 16 就 OOM。代码看起来也没问题model.eval()加了torch.no_grad()也加了但显存还是居高不下。问题往往不在模型本身而在于这两个 API 的语义被混为一谈或者只加了一个。model.eval()和torch.no_grad()是 PyTorch 推理阶段最常被同时提及、却最容易被误解的一对组合。它们作用在不同的层面前者改变的是网络层的前向行为后者改变的是 autograd 引擎的记账行为。一个管“算得对不对”一个管“算得省不省”。把这两件事拆开理解你才能知道为什么四种组合下显存和耗时差异会那么大。这篇文章面向的是已经能把模型跑起来、但想把推理服务压到更省显存、更低延迟的开发者。场景选图像分类模型部署因为 ResNet、ViT 这类模型对 BatchNorm 和 Dropout 的依赖非常典型四种组合的差异会被放大得很明显。我会给出可直接复制的 benchmark 脚本、显存统计代码、验证步骤最后说明如何通过 TaoToken 统一 Key 与 API 通道把本地推理服务接到端到端压测里。先给结论方便你对照自己的代码model.eval()负责把 Dropout 关掉、把 BatchNorm 切到用 running_mean/running_vartorch.no_grad()负责停止构建计算图从而省掉中间激活的梯度存储。两者互不影响同时使用才是推理的正确姿势。只加model.eval()显存省不下来只加torch.no_grad()BatchNorm 还在更新统计量结果可能飘。2. model.eval() 到底改了什么Dropout 与 BatchNorm 的行为切换2.1 Dropout 在 eval 下的真实行为训练时 Dropout 会按概率 p 随机把一部分激活置零并把保留的激活除以保留概率做缩放保证期望不变。到了 evalDropout 不再随机丢弃而是让所有激活单元通过。这里有个常见疑问训练时明明屏蔽了一些神经元推理时全放行预测还准吗用个类比训练像限定你每次只能翻一份资料逼你学会不依赖单一来源考试时所有资料都摊开但你心里清楚每份资料的权重。eval 下所有激活都通过但各神经元的输出会按训练时的保留比例做等效缩放所以整体期望是一致的。这也是为什么model.eval()必须加否则推理结果会带随机性同一个输入两次跑出来的 logits 都不一样。2.2 BatchNorm 在 eval 下用的是 running 统计量BatchNorm 在 train 模式会用当前 batch 的 mean 和 var 做归一化并更新 running_mean、running_var。eval 模式则停止更新直接用训练阶段累积的 running 统计量。这一点对推理至关重要如果推理时 batch size 很小比如 1用当前 batch 的统计量会非常不稳定结果直接崩掉。所以model.eval()不是可选项而是保证推理可复现、结果稳定的前提。它不影响梯度计算本身梯度该建还是建只是前向行为变了。真正省显存的那一步得靠torch.no_grad()。2.3 一个容易踩的坑忘了切回 train我见过不少代码在验证集上跑完model.eval()接着继续训练却忘了model.train()结果后面几个 epoch 的 BatchNorm 统计量全乱loss 曲线突然抖。建议把 eval 和 train 的切换封装成上下文管理器或者至少在验证函数入口和出口成对写。下面这段可以直接用import torch import torch.nn as nn from contextlib import contextmanager contextmanager def eval_mode(model): was_training model.training model.eval() try: yield model finally: model.train(was_training)这样无论验证过程中是否抛异常模型状态都能恢复避免污染后续训练。3. torch.no_grad() 与四种组合的可复制 benchmark 配置3.1 no_grad 省的是什么torch.no_grad()关闭的是 autograd 的图构建。默认情况下每个需要梯度的张量运算都会记录操作形成计算图中间激活值要保留下来供反向传播用。推理根本不需要反向这些中间激活就是纯浪费。关掉之后前向只算结果不留图显存和算力都能省下来。注意它不影响 Dropout 和 BatchNorm 的行为那是model.eval()的职责。两者正交可以自由组合。3.2 四种组合的 benchmark 脚本下面这段脚本对同一模型跑四种组合统计显存峰值和单批耗时。模型用 torchvision 的 resnet50输入固定尺寸保证可比。import torch import torch.nn as nn import time from torchvision.models import resnet50 def build_model(): model resnet50(weightsNone) model.fc nn.Linear(model.fc.in_features, 10) return model.cuda() def measure(model, use_eval, use_no_grad, batch_size32, iters20): x torch.randn(batch_size, 3, 224, 224).cuda() if use_eval: model.eval() else: model.train() torch.cuda.empty_cache() torch.cuda.reset_peak_memory_stats() # warmup for _ in range(3): if use_no_grad: with torch.no_grad(): _ model(x) else: _ model(x) torch.cuda.synchronize() start time.time() for _ in range(iters): if use_no_grad: with torch.no_grad(): out model(x) else: out model(x) torch.cuda.synchronize() elapsed (time.time() - start) / iters peak_mem torch.cuda.max_memory_allocated() / 1024**2 return elapsed * 1000, peak_mem if __name__ __main__: model build_model() configs [ (train grad, False, False), (train no_grad, False, True), (eval grad, True, False), (eval no_grad, True, True), ] for name, use_eval, use_no_grad in configs: ms, mem measure(model, use_eval, use_no_grad) print(f{name:20s} | {ms:8.2f} ms/batch | peak {mem:8.1f} MB)跑下来典型结果不同卡会有差异但趋势一致eval no_grad显存峰值最低、耗时最短train grad最费eval grad显存和train grad接近因为图还在建train no_grad省显存但 BatchNorm 还在更新结果不可用于推理。3.3 用 JSON 固化推理配置如果你要把推理参数交给服务化框架建议用 JSON 显式声明避免运行时靠默认值猜{ model_name: resnet50, input_size: [3, 224, 224], batch_size: 32, eval_mode: true, no_grad: true, device: cuda, dtype: float32 }这份配置里eval_mode和no_grad分开写就是为了提醒自己这是两件事。很多线上事故就是有人把no_grad当成eval的替代品结果 BatchNorm 行为不对指标掉点还找不到原因。4. 验证请求与成功结果显存统计与端到端压测4.1 显存统计要看得细torch.cuda.max_memory_allocated()看的是 PyTorch 分配器记录的峰值torch.cuda.memory_reserved()看的是缓存池保留量。排查 OOM 时两个都要看因为有时候 allocated 不高但 reserved 涨得厉害是碎片问题。下面这段可以打印更细的分布def report_memory(tag): allocated torch.cuda.memory_allocated() / 1024**2 reserved torch.cuda.memory_reserved() / 1024**2 peak torch.cuda.max_memory_allocated() / 1024**2 print(f[{tag}] allocated{allocated:.1f}MB reserved{reserved:.1f}MB peak{peak:.1f}MB)在 benchmark 每个配置前后各调一次就能看到显存是否被正确释放。如果eval no_grad跑完 reserved 还很高说明缓存池没回收可以手动torch.cuda.empty_cache()但生产环境不建议频繁调用会拖慢。4.2 用 TaoToken 统一通道做端到端压测本地 benchmark 只能说明单机单卡的表现真实服务还要看请求链路。我习惯把推理服务包一层 HTTP 接口然后用统一的 Key 和 API 通道去压。TaoToken 在这里的作用是把模型调用、Key 管理、额度统计收敛到一个入口省得每个服务各配一套。接入时 Base URL 用https://taotoken.net/apiKey 在控制台生成Model ID 按你实际部署的模型名填。三件套缺一不可尤其是 Model ID写错会直接报模型不存在。配置片段如下{ base_url: https://taotoken.net/api, api_key: sk-你的Key, model_id: resnet50-infer }压测脚本可以用 requests 并发打观察 P99 延迟和错误率。重点看两件事一是eval no_grad下延迟是否稳定二是并发升高时显存是否线性增长。如果显存随并发涨多半是每个请求都新建了计算图检查是不是漏了no_grad。4.3 成功结果的判断标准一次合格的推理压测应该满足同一输入多次请求输出一致证明 eval 生效、显存峰值不随请求数累积证明 no_grad 生效、P99 延迟在可接受范围。三条都过才算真正把这两个 API 用对了。5. 本篇常见报错排查401、local proxy failed 与 reading choices5.1 401 Unauthorized最常见的是 Key 没带或带错。检查请求头是不是Authorization: Bearer sk-xxx注意 Bearer 后面有空格。如果用的是环境变量确认变量名和读取代码一致。还有一种情况是 Key 被禁用或额度耗尽去控制台看状态。5.2 local proxy failed这个报错通常出现在本地服务转发请求时代理配置指向了不可达地址。排查顺序先确认服务监听端口再确认转发目标 URL 拼写最后看防火墙。注意不要在任何配置里写来路不明的转发地址统一走https://taotoken.net/api这类明确入口避免链路不可控。5.3 reading choices 相关报错如果返回体解析时报reading choices或类似字段缺失说明响应结构和你预期的不一致。先打印原始 response.text 看真实返回再对照文档确认字段路径。常见原因是 Model ID 写错导致返回了错误对象或者请求体里 messages 格式不对。把原始返回打出来比猜快得多。5.4 OAuth 与鉴权混淆有些框架默认走 OAuth 流程和 API Key 是两套东西。如果你只用 Key就别开 OAuth 相关开关否则会卡在跳转。确认鉴权方式单一减少排查面。5.5 排查清单遇到问题按这个顺序走Key 是否存在且未过期 → Base URL 是否正确 → Model ID 是否匹配 → 请求体字段是否完整 → 原始返回是什么。五步走完绝大多数报错都能定位。6. 把推理配置沉淀成可复用资产推理阶段的优化不是加两个 API 就完事而是要把配置、脚本、压测流程沉淀下来。我的做法是benchmark 脚本进仓库JSON 配置进版本管理压测结果按版本归档。这样每次模型更新跑一遍就能知道显存和延迟有没有退化。TaoToken 的接入文档里有完整的鉴权和调用示例模型对话入口可以用来快速验证通道是否通Coding Plan 适合长期跑 Agent 类任务API Keys 页面管理你的凭证。把这些入口固定下来团队里谁接手都能快速复现。最后留一个实用技巧在推理服务启动时打印一次model.training和torch.is_grad_enabled()确认状态符合预期。这行日志能帮你省掉很多“为什么结果不对”的排查时间。
返回列表