PyTorch 自动求导(Autograd)
🤖 什么是自动求导?
自动求导(Autograd)是 PyTorch 最核心的特性之一,它自动计算神经网络中所有参数的梯度,让你无需手动推导和实现反向传播。PyTorch 使用动态计算图——每次前向传播都构建一张新图,这意味着你可以使用 Python 原生的控制流(if/for/while),模型结构可以在每次迭代中改变。
🎯 核心概念:requires_grad
importtorch# 创建可求导的张量x=torch.tensor([1.0,2.0,3.0],requires_grad=True)print(x.requires_grad)# True# 或者事后设置x=torch.tensor([1.0,2.0,3.0])x.requires_grad_(True)# 原地修改(注意尾部 _)requires_grad=True告诉 PyTorch:这个张量的所有运算都要被追踪,以便后续计算梯度。
📐 梯度计算基本流程
# 1. 定义变量(需要梯度)x=torch.tensor([2.0,3.0],requires_grad=True)# 2. 定义计算(构建计算图)y=x**2+3*x+1# y = x² + 3x + 1# dy/dx = 2x + 3# 3. 调用 backward() 计算梯度y.sum().backward()# y 是向量,需要转为标量# 4. 查看梯度print(x.grad)# tensor([7., 9.]) ← 2*2+3=7, 2*3+3=9 ✓为什么需要.sum()再backward()?
backward()只能对标量调用。如果y是向量/矩阵,需要先将其转为标量(求和、求均值等)。
x=torch.randn(3,requires_grad=True)y=x*2# ❌ y.backward() ← 报错!y 是向量# ✅ y.sum().backward() ← 正确:先转为标量# ✅ y.mean().backward() ← 也可以取均值🧮 梯度计算原理
x=torch.tensor([1.0,2.0,3.0],requires_grad=True)w=torch.tensor([0.5,0.3,0.2],requires_grad=True)b=torch.tensor(0.1,requires_grad=True)# 前向传播z=(x*w).sum()+b# 标量输出# 反向传播(自动计算 dz/dx, dz/dw, dz/db)z.backward()print(f'dz/dx:{x.grad}')# w = [0.5, 0.3, 0.2]print(f'dz/dw:{w.grad}')# x = [1.0, 2.0, 3.0]print(f'dz/db:{b.grad}')# 1.0🔄 梯度累积与清零 ⭐
PyTorch 默认累加梯度(而非覆盖),这是为 RNN 等多步反向传播设计的。
x=torch.tensor([2.0],requires_grad=True)foriinrange(3):y=x**2y.backward()print(f'Step{i}: grad ={x.grad}')# Step 0: grad = [4.] ← 2*2# Step 1: grad = [8.] ← 累加!4 + 4# Step 2: grad = [12.] ← 再累加!8 + 4# 正确做法:每次迭代清空梯度x=torch.tensor([2.0],requires_grad=True)foriinrange(3):y=x**2y.backward()print(f'Step{i}: grad ={x.grad}')x.grad.zero_()# ← 清零!在实际训练中,使用optimizer.zero_grad()一次性清空所有参数梯度。
🚫 禁用梯度追踪
torch.no_grad()— 禁用梯度计算
在推理、评估时使用,减少内存占用、加速计算。
x=torch.randn(10,requires_grad=True)# 方式一: with 语句(推荐)withtorch.no_grad():y=x*2print(y.requires_grad)# False# 方式二: 装饰器@torch.no_grad()defevaluate(model,data):returnmodel(data).detach()— 从计算图中分离
创建一个共享数据但不参与梯度计算的新张量。
x=torch.tensor([2.0],requires_grad=True)y=x**2# y = 4z=y.detach()# z 是 4,但不与 x 关联w=z**2# 与 x 无关,不会反向传播到 xset_grad_enabled()— 条件控制
withtorch.set_grad_enabled(is_train):output=model(input)loss=loss_fn(output,target)⛓️ 阻止对部分参数求导
# 方式一: requires_grad=Falsew=torch.randn(3,5,requires_grad=False)# 不更新b=torch.randn(5,requires_grad=True)# 更新# 方式二: 冻结模型层forparaminmodel.features.parameters():param.requires_grad=False# 冻结 backbone# 方式三: optimizer 只传部分参数optimizer=optim.Adam(filter(lambdap:p.requires_grad,model.parameters()),lr=1e-4)📊 梯度信息查看与调试
x=torch.tensor([[1.0,2.0],[3.0,4.0]],requires_grad=True)y=(x**2).sum()y.backward()print(x.grad)# tensor([[2., 4.], [6., 8.]])# 梯度相关属性print(x.grad_fn)# None(叶子节点)print(y.grad_fn)# <SumBackward0> — y 是由 sum 得到的# 查看计算图print(x.is_leaf)# True(用户直接创建的)叶子节点(leaf node):用户直接创建(而非运算结果)的张量。只有叶子节点的.grad会在backward()后被填充。
🔧 retain_graph — 保留计算图
默认backward()后计算图被释放。如果需要对同一个输出多次调用backward():
x=torch.tensor([2.0],requires_grad=True)y=x**3# y = 8y.backward(retain_graph=True)print(x.grad)# 12 ← 3*2² = 12y.backward(retain_graph=True)print(x.grad)# 24 ← 累加到 12+12 = 24(注意梯度累加!)⚠️ 绝大多数情况下不需要
retain_graph=True。只在多任务学习、GAN 训练等少数场景用到。
🧪 高阶梯度
PyTorch 支持计算梯度的梯度(二阶导数)。
x=torch.tensor([2.0],requires_grad=True)y=x**3# y = x³# 一阶导: dy/dx = 3x² = 12grad1=torch.autograd.grad(y,x,create_graph=True)[0]print(grad1)# tensor([12.])# 二阶导: d²y/dx² = 6x = 12grad2=torch.autograd.grad(grad1,x)[0]print(grad2)# tensor([12.])使用create_graph=True保留梯度计算图,以便计算更高阶导数。
⚠️ 常见错误与解决
| 错误 | 原因 | 解决 |
|---|---|---|
backward()报错 | 输出不是标量 | 先.sum()或.mean() |
grad为None | 没设requires_grad=True | 创建时设置或调用.requires_grad_() |
grad值不对 | 忘记清零,梯度累加 | 每次迭代前optimizer.zero_grad() |
| 内存不足 | 保留了不必要的计算图 | 用torch.no_grad()或.detach() |
inplace操作报错 | 修改了需要梯度的叶子节点 | 避免对需要梯度的张量做原地操作 |
📝 速查表
| 需求 | 代码 |
|---|---|
| 启用梯度 | x = torch.tensor([1.], requires_grad=True) |
| 反向传播 | loss.backward() |
| 查看梯度 | x.grad |
| 清零梯度 | x.grad.zero_() |
| 禁用梯度 | with torch.no_grad(): |
| 分离张量 | x.detach() |
| 冻结参数 | param.requires_grad = False |
| 获取梯度值 | torch.autograd.grad(loss, x) |
| 保留计算图 | loss.backward(retain_graph=True) |
| 创建高阶图 | torch.autograd.grad(y, x, create_graph=True) |
[[pytorch-总览|← 返回总览]]