1. 训练循环里最容易踩的坑:模式没切对,指标全白费
如果你写过 PyTorch 训练脚本,大概率见过这几个调用:model.train()、model.eval()、torch.no_grad()、detach()。它们看起来都是「一行代码」,但作用的对象完全不同:前两个改的是模型内部层的行为状态,后两个管的是计算图和梯度记录。混着用,轻则验证集指标忽高忽低,重则显存爆掉、梯度算错,训练半天不收敛。
我见过最常见的翻车场景是这样的:训练循环里忘了写model.eval(),验证时 Dropout 还在随机丢神经元,BatchNorm 还在用当前 batch 的统计量,于是同一个模型跑两遍验证集得到两个不同的准确率,你还以为是数据有问题。另一种是推理时没套torch.no_grad(),明明只是前向计算,却把整个计算图都建起来了,显存占用翻倍,batch 稍微大一点就 OOM。
这篇就围绕这四个调用,给你一套可以直接复制的训练/验证循环骨架,再配上「打印 requires_grad 和 grad_fn」的验证动作,让你亲眼确认梯度状态对不对。适合刚接触 PyTorch 的开发者,也适合写了很久但一直靠感觉写循环的人。下面所有代码都可以直接跑,不需要额外数据集,用随机张量就能验证行为。
2. 先把 TaoToken 配好:让训练脚本能调到大模型做辅助
写训练代码时,经常需要让模型帮忙解释报错、生成数据增强脚本,或者对比不同实现。这时候一个稳定的 API 入口能省不少事。TaoToken 提供统一的模型调用入口,官网是 https://taotoken.net/?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite&utm_content= ,API 地址是 https://taotoken.net/api ,注意 API 地址后面不加 UTM 参数。
配置方式很简单,把 API Key 写进环境变量,避免硬编码进代码:
export TAOTOKEN_API_KEY="你的key" export TAOTOKEN_BASE_URL="https://taotoken.net/api"然后在 Python 里读取:
import os api_key = os.environ["TAOTOKEN_API_KEY"] base_url = os.environ["TAOTOKEN_BASE_URL"] print(base_url)如果你还没拿到 Key,可以去控制台创建:https://taotoken.net/console?utm_source=taotoken_aicg_blog_end&utm_content=console&utm_campaign=rewrite ,Key 管理页面在 https://taotoken.net/api-keys?utm_source=taotoken_aicg_blog_end&utm_content=api-keys&utm_campaign=rewrite 。接入文档在 https://taotoken.net/doc?utm_source=taotoken_aicg_blog_end&utm_content=doc&utm_campaign=rewrite ,遇到参数问题先翻文档比瞎试快。
注意:API Key 只放在环境变量或本地配置文件里,不要提交到 Git 仓库。训练脚本里也不要打印完整 Key。
配好之后,你可以在训练脚本旁边写个小工具函数,把报错信息丢给模型对话接口问原因,模型对话入口是 https://taotoken.net/model-chat?utm_source=taotoken_aicg_blog_end&utm_content=model-chat&utm_campaign=rewrite 。这样调试训练循环时效率会高很多。
3. 可复制配置:训练循环与验证循环骨架
3.1 四个调用的职责划分
先把概念理清楚,后面写代码才不会乱。
model.train()把模型的training属性设为 True。它影响的是 Dropout 和 BatchNorm 这类「训练/推理行为不同」的层。Dropout 在训练时按概率丢弃神经元,推理时全部保留;BatchNorm 在训练时用当前 batch 的均值和方差,并更新 running 统计量,推理时用 running 统计量。
model.eval()把training设为 False,上面两类层切换到推理行为。它不涉及梯度,也不释放显存。
torch.no_grad()是上下文管理器,进入后所有计算不构建计算图,requires_grad即使为 True 的张量,运算结果也不会带grad_fn。它省的是显存和计算,常用于验证和推理。
detach()是张量方法,从计算图里「切」出一个新张量,新张量requires_grad=False,和原张量共享数据。常用于把 loss 或中间结果拿出来做日志、算指标,防止它们把梯度图拖住。
一句话对照:
| 调用 | 作用对象 | 影响梯度 | 影响层行为 | 典型位置 |
|---|---|---|---|---|
| model.train() | 模型 | 否 | 是 | 训练循环开头 |
| model.eval() | 模型 | 否 | 是 | 验证/推理开头 |
| torch.no_grad() | 计算图 | 是 | 否 | 验证/推理包裹 |
| detach() | 单个张量 | 是 | 否 | 日志/指标计算 |
3.2 训练循环骨架
import torch import torch.nn as nn def train_one_epoch(model, loader, optimizer, criterion, device): model.train() # 关键:切到训练模式 total_loss = 0.0 for x, y in loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() logits = model(x) loss = criterion(logits, y) loss.backward() optimizer.step() total_loss += loss.item() # item() 已脱离计算图 return total_loss / len(loader)这里loss.item()本身就返回 Python 标量,不需要 detach。但如果你要累积一个 tensor 形式的 loss,就必须 detach:
running = torch.zeros(1, device=device) running += loss.detach() # 不 detach 会把整个图累积起来3.3 验证循环骨架
@torch.no_grad() def evaluate(model, loader, criterion, device): model.eval() # 关键:切到推理模式 total_loss = 0.0 correct = 0 total = 0 for x, y in loader: x, y = x.to(device), y.to(device) logits = model(x) loss = criterion(logits, y) total_loss += loss.item() pred = logits.argmax(dim=1) correct += (pred == y).sum().item() total += y.size(0) acc = correct / total return total_loss / len(loader), acc注意@torch.no_grad()装饰器和model.eval()是两件事,都要写。前者管梯度,后者管层行为。少任何一个都会出问题。
3.4 detach 的正确使用位置
验证时算指标,如果不用no_grad,至少要对参与累积的张量 detach:
with torch.no_grad(): logits = model(x) pred = logits.argmax(dim=1) correct += (pred == y).sum().item()如果忘了no_grad,pred == y的结果是 bool tensor,.sum()会带 grad_fn,.item()虽然能取值,但中间图已经建好了,显存白占。所以要么整体no_grad,要么对每个要累积的量 detach。
4. 验证请求:打印 requires_grad 与 grad_fn 确认状态
光看代码不够,跑一遍打印出来才踏实。下面这段可以直接复制运行。
import torch import torch.nn as nn class Net(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(4, 2) self.drop = nn.Dropout(0.5) self.bn = nn.BatchNorm1d(2) def forward(self, x): x = self.fc(x) x = self.bn(x) return self.drop(x) model = Net() x = torch.randn(8, 4) # 训练模式 model.train() out_train = model(x) print("train mode, training =", model.training) print("out_train.requires_grad =", out_train.requires_grad) print("out_train.grad_fn =", out_train.grad_fn) # 推理模式 + no_grad model.eval() with torch.no_grad(): out_eval = model(x) print("eval mode, training =", model.training) print("out_eval.requires_grad =", out_eval.requires_grad) print("out_eval.grad_fn =", out_eval.grad_fn) # detach 验证 a = torch.tensor([1.1], requires_grad=True) b = a.detach() print("a.requires_grad =", a.requires_grad) print("b.requires_grad =", b.requires_grad) print("b.grad_fn =", b.grad_fn)预期输出:
train mode, training = True out_train.requires_grad = True out_train.grad_fn = <AddmmBackward0 ...> eval mode, training = False out_eval.requires_grad = False out_eval.grad_fn = None a.requires_grad = True b.requires_grad = False b.grad_fn = None看到out_eval.grad_fn = None就说明no_grad生效了。如果这里打印出非 None,说明你的no_grad没包住前向,或者模型里有地方偷偷建了图。
再补一个 detach 的对比实验,确认它只影响新张量:
x = torch.randn(3, 2, requires_grad=True) w = torch.tensor([1.1, 2.2]) b = torch.ones(3) z1 = torch.matmul(x, w) + b print("z1.requires_grad =", z1.requires_grad) # True with torch.no_grad(): z2 = torch.matmul(x, w) + b print("z2.requires_grad =", z2.requires_grad) # False z3 = z1.detach() print("z3.requires_grad =", z3.requires_grad) # False print("x.requires_grad =", x.requires_grad) # True,原张量不受影响5. 本篇常见错排查
5.1 验证指标每次都不一样
现象:同一个模型、同一份验证集,跑两次准确率差好几个点。
原因:验证循环里没写model.eval(),Dropout 还在随机丢弃,BatchNorm 还在用当前 batch 统计量。
排查:在验证循环开头打印model.training,应该是 False。如果打印 True,就是漏了model.eval()。
5.2 显存越跑越大,最后 OOM
现象:训练几个 epoch 后显存持续上涨。
原因:把带梯度的 tensor 累积进了列表或变量。比如total_loss += loss而不是loss.item()或loss.detach()。
排查:检查所有累积操作,凡是 tensor 相加的地方,确认右边是否 detach 或用了 item()。
5.3 报错 "element 0 of tensors does not require grad"
现象:调用loss.backward()时报这个错。
原因:loss 的requires_grad是 False。常见于验证阶段误调 backward,或者模型参数被冻结后没解冻,或者输入张量被 detach 过。
排查:打印loss.requires_grad,确认是 True 再 backward。验证阶段本来就不该 backward。
5.4 no_grad 里又开了 requires_grad
现象:明明套了torch.no_grad(),结果张量还是有 grad_fn。
原因:在no_grad块里手动把某个张量的requires_grad设回 True,或者调用了torch.enable_grad()。
排查:搜索代码里有没有requires_grad_(True)或enable_grad,确认它们不在推理路径上。
5.5 detach 后还想 backward
现象:对 detach 出来的张量调 backward,报错或梯度为 None。
原因:detach 就是切断计算图,切断了自然回不去。如果你需要保留梯度又要取值,用.clone()而不是.detach(),或者只在日志场景用 detach。
排查:确认 detach 的使用场景是「只读不反传」,需要反传的路径不要 detach。
5.6 BatchNorm 在 batch size 为 1 时训练报错
现象:Expected more than 1 value per channel when training。
原因:BatchNorm 在训练模式下需要 batch 内多个样本算方差,batch size 为 1 时无法计算。
排查:要么调大 batch size,要么在model.train()前对 BN 层单独设eval(),要么改用 GroupNorm。这个和model.train()的切换直接相关,别在推理时误开训练模式。
6. 把模式切换写进模板,长期编码更省心
上面这套骨架,建议直接固化成一个训练模板文件,每次新项目复制过去改模型和数据集就行。模板里把model.train()、model.eval()、torch.no_grad()、detach()的位置都标好注释,减少遗漏。
如果你经常写训练脚本、Agent 工具链或者需要反复调试模型调用,可以考虑用 Coding Plan 把常用代码片段和 API 调用统一管理,入口在 https://taotoken.net/coding-plan?utm_source=taotoken_aicg_blog_end&utm_content=coding-plan&utm_campaign=rewrite 。配合 API Keys 页面 https://taotoken.net/api-keys?utm_source=taotoken_aicg_blog_end&utm_content=api-keys&utm_campaign=rewrite 管理密钥,接入文档 https://taotoken.net/doc?utm_source=taotoken_aicg_blog_end&utm_content=doc&utm_campaign=rewrite 查参数,基本能覆盖日常开发。
最后留一个我常用的自检习惯:每个 epoch 结束后,打印一次model.training和最近一个 batch 的loss.grad_fn,确认训练模式是 True、梯度图正常;验证结束后再打印一次,确认是 False、grad_fn 为 None。两行 print,能挡掉大部分模式切换的坑。