PyTorch中含Softmax与log2的函数二阶导数全零问题排查
问题描述
我尝试使用PyTorch计算函数t相对于张量a的二阶导数(Hessian矩阵),初始代码如下:
import torch torch.manual_seed(0) a = torch.randint(0, 10, (10,), dtype=float, requires_grad=True) b, *_ = a.sort(descending=True) c = (b.unsqueeze(0) - a.unsqueeze(1)).abs().neg() d = c.softmax(0).matmul(torch.arange(c.size(0), dtype=c.dtype)) e = torch.randint(0, 3, (10,), dtype=float) t = torch.sum(e * torch.log2(d + 1)) grad, *_ = torch.autograd.grad(t, a, create_graph=True) hess, *_ = torch.autograd.grad(grad, a, torch.ones_like(a)) print(hess)
我期望代码能计算t对a的逐元素二阶导数,但返回的Hessian向量全为0。这让我困惑,因为函数包含softmax和对数操作,本应产生非零二阶导数。
为进一步排查,我尝试了另一种方法:
import torch torch.manual_seed(0) a = torch.randint(0, 10, (10,), dtype=float, requires_grad=True) b, *_ = a.sort(descending=True) c = (b.unsqueeze(0) - a.unsqueeze(1)).abs().neg() d = c.softmax(0).matmul(torch.arange(c.size(0), dtype=c.dtype)) e = torch.randint(0, 3, (10,), dtype=float) t = torch.sum(e * torch.log2(d + 1)) t.backward(create_graph=True) grad = a.grad.clone() hess = torch.zeros_like(a) for i in range(len(a)): a.grad.zero_() grad[i].backward(retain_graph=True) hess[i] = a.grad[i] print(hess)
第二种方法中,梯度grad与第一种方法一致,但Hessianhess结果不同。两种方法存在差异,但我不清楚原因。
问题:
- 为何第一种方法返回全零的Hessian?
- 哪种方法能正确计算该场景下的二阶导数?
- 若均不正确,正确的Hessian计算方法是什么?
解答
1. 第一种方法返回全零Hessian的原因
第一种方法中,torch.autograd.grad(grad, a, torch.ones_like(a))计算的是梯度向量与全1向量的点积(即所有梯度元素之和)对a的梯度,对应Hessian矩阵每一列元素的和,而非你预期的逐元素二阶导数(Hessian矩阵的对角线)。
在你固定seed=0的输入场景下,a存在重复元素,结合sort、softmax的计算逻辑,恰好出现了Hessian矩阵每一列元素之和为0的特殊情况,因此返回全零向量。这并非二阶导数本身全零,而是计算目标与预期不符导致的结果。
2. 哪种方法能正确计算逐元素二阶导数
第二种方法是正确的,它通过逐个对梯度的第i个元素求导,得到t对a[i]的二阶导数(即Hessian矩阵的第i个对角线元素),完全匹配你期望的“逐元素二阶导数”需求。
需要注意的是,该方法仅计算了Hessian的对角线元素,若你需要交叉二阶偏导数(完整Hessian矩阵),则需要修改代码保存整个a.grad而非仅a.grad[i]。
3. 完整Hessian矩阵的正确计算方法
若需要计算包含所有交叉偏导数的完整Hessian矩阵,可以修改第二种方法的循环逻辑,保存每个grad[i]对所有a[j]的导数:
import torch torch.manual_seed(0) a = torch.randint(0, 10, (10,), dtype=torch.float64, requires_grad=True) b, *_ = a.sort(descending=True) c = (b.unsqueeze(0) - a.unsqueeze(1)).abs().neg() d = c.softmax(0).matmul(torch.arange(c.size(0), dtype=c.dtype)) e = torch.randint(0, 3, (10,), dtype=torch.float64) t = torch.sum(e * torch.log2(d + 1)) # 计算一阶梯度并保留计算图 t.backward(create_graph=True) grad = a.grad.clone() # 初始化完整Hessian矩阵 hessian = torch.zeros((len(a), len(a)), dtype=torch.float64) for i in range(len(a)): a.grad.zero_() # 对第i个梯度元素求导,保留计算图 grad[i].backward(retain_graph=True) # 保存第i行的所有Hessian元素 hessian[i] = a.grad.clone() print(hessian)
这段代码会生成10×10的矩阵,其中hessian[i][j]表示t对a[i]和a[j]的二阶混合偏导数。
内容的提问来源于stack exchange,提问作者Ray Bern

