PyTorch clip_grad_norm_函数未生效?求梯度裁剪原理与正确示例
为什么你的PyTorch梯度裁剪没生效?
你犯了一个核心错误:torch.nn.utils.clip_grad_norm_是用来裁剪张量的梯度(.grad属性),不是直接操作普通张量。你代码里的v_1只是普通张量,没有梯度信息,所以函数根本没执行任何裁剪操作。
正确的梯度裁剪示例
要让clip_grad_norm_生效,必须先让张量产生梯度,再对梯度进行裁剪:
import torch # 1. 创建需要计算梯度的张量,必须设置requires_grad=True v = torch.rand(5)*1000 v.requires_grad = True # 2. 模拟正向传播,构造一个能产生大梯度的loss(这里用元素平方和,导数为2*v) loss = (v ** 2).sum() # 3. 反向传播,生成梯度(存储在v.grad中) loss.backward() print("原始梯度:", v.grad) original_norm = torch.norm(v.grad, p=2) print("原始梯度L2范数:", original_norm.item()) # 4. 执行梯度裁剪:传入包含目标张量的列表,函数会自动处理其.grad属性 torch.nn.utils.clip_grad_norm_([v], max_norm=1.0, norm_type=2) print("\n裁剪后的梯度:", v.grad) clipped_norm = torch.norm(v.grad, p=2) print("裁剪后的梯度L2范数:", clipped_norm.item())
原理说明
clip_grad_norm_的工作逻辑:
- 计算所有传入张量的梯度的总L2范数(你指定
norm_type=2) - 如果总范数大于
max_norm,就对所有梯度执行缩放:梯度 = 梯度 * (max_norm / 总范数) - 最终所有梯度的总范数会被限制为
max_norm
你之前的代码直接传入无梯度的普通张量,函数内部检测不到可裁剪的梯度,所以没有任何变化。
内容的提问来源于stack exchange,提问作者max_max_mir
相关产品推荐
相关产品推荐

