PyTorch中正确计算子矩阵梯度的方法
报错原因
你是在损失计算完成后才执行A[:1, :]切片操作,得到的子张量从未参与损失对应的前向计算流程,不存在于PyTorch记录的计算图中,自动微分引擎无法识别它和损失的关联,因此抛出对应错误。
损失相对于A的子矩阵的梯度,在数学上完全等价于损失相对于完整A的梯度对应位置的切片,因此最直接高效的解法如下:
正确实现代码
import torch from torch import nn A = torch.rand((3, 5), requires_grad=True) target = torch.zeros((3, 5), requires_grad=True) loss = nn.functional.mse_loss(A, target) # 先计算完整梯度 full_grad = torch.autograd.grad(loss, A)[0] print(f'full grad shape = {full_grad.shape}') # 对完整梯度做切片,直接得到子矩阵对应的梯度 sub_grad = full_grad[:1, :] print(f'grad shape = {sub_grad.shape}')
这种实现没有额外计算开销,结果完全符合数学定义,适用于绝大多数场景。
内容的提问来源于stack exchange,提问作者Tengerye
相关产品推荐
相关产品推荐

