如何在不丢失梯度的情况下对Tensor进行掩码操作?
解决方法:切断掩码的梯度传播,保留原张量梯度
你的问题出在torch.no_grad()会禁用整个代码块的梯度追踪,导致结果张量b也失去了梯度属性。正确的做法是只让掩码张量不参与梯度传播,同时保留原张量a的梯度追踪,用mask.detach()就能实现。
示例代码
import torch a = torch.randn(1, 3, requires_grad=True) print('a: ', a) >>> a: tensor([[0.0200, 1.0020, -4.2000]], requires_grad=True) # 实际场景中掩码带有梯度 mask = torch.zeros_like(a, requires_grad=True) mask[0][0] = 1 # 用detach()切断mask的梯度传播 b = a * mask.detach() print('b: ', b) >>> b: tensor([[0.0200, 0.0000, -0.0000]], grad_fn=<MulBackward0>) print('b.requires_grad: ', b.requires_grad) >>> b.requires_grad: True
为什么这么做?
mask.detach()返回一个和原mask数据完全一致,但不记录梯度的张量。和a相乘时,梯度只会流向a,不会传递给原mask张量。- 而
torch.no_grad()会让整个代码块里的所有运算都不追踪梯度,导致b的requires_grad直接变成False,完全丢失了梯度信息,没法继续后续反向传播。
验证梯度是否符合预期
可以通过反向传播测试:
b.sum().backward() print('a的梯度: ', a.grad) >>> a的梯度: tensor([[1., 0., 0.]]) # 梯度正确传到a的对应位置 print('mask的梯度: ', mask.grad) >>> mask的梯度: None # mask没有得到梯度,符合需求
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

