PyTorch计算两层线性网络Hessian条件数与理论值不符求助
两层线性网络Hessian条件数计算差异排查
问题背景
训练一个两层线性网络,求解优化问题:
$$\min_{W_1, W_2} \frac{1}{2}\lVert Y-W_2W_1X\rVert_F^2$$
其中$X,Y$为$\mathbb{R}^{5\times 5}$矩阵,$W_1,W_2$是宽度为100的参数。关注迭代接近最小值时损失关于$W_2$的Hessian条件数,两种计算方法结果差异显著:
- 方法一:PyTorch数值计算Hessian,训练至相对损失达$10^{-20}$,条件数约174
- 方法二:全局最小值处Hessian解析形式$H=(W_1XX\topW_1\top)\otimes I_5$,代入训练后$W_1$计算得条件数为5345
复现代码
import torch import torch.nn as nn import numpy as np from scipy.stats import ortho_group class NN_linear(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(NN_linear, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size, bias=False) self.fc2 = nn.Linear(hidden_size, output_size, bias=False) # Initialize weights with custom numpy matrices np.random.seed(0) self.fc1.weight = nn.Parameter(torch.from_numpy(np.random.normal(0, 1, size= (input_size, hidden_size)).T).double()) self.fc2.weight = nn.Parameter(torch.from_numpy(np.random.normal(0, 1, size= (hidden_size, output_size)).T).double()) def forward(self, x): x = self.fc1(x) x = self.fc2(x) return x # data generation kappa = 100 input_dim = 5 output_dim = 5 width = 100 N = 5 np.random.seed(123) theta = np.random.normal(0, 1, size=(input_dim, output_dim)).astype(np.float64) u = ortho_group.rvs(dim=input_dim).astype(np.float64) v = ortho_group.rvs(dim=output_dim).astype(np.float64) s = np.diag(np.linspace(start=1, stop=np.sqrt(kappa), endpoint=True, num=N)) x = u.dot(s.dot(v)) y = x.dot(theta) dataset = torch.utils.data.TensorDataset(torch.tensor(x, dtype=torch.double), torch.tensor(y, dtype=torch.double)) dataloader = torch.utils.data.DataLoader( dataset=dataset, batch_size=N, shuffle=False) optimal_lr = 0.00023400934009340095 lr = optimal_lr # 补充缺失的lr赋值 # training net = NN_linear(input_dim, width, output_dim) optimizer = torch.optim.SGD(net.parameters(), lr=lr) epoch = 0 loss_initial = None while epoch == 0 or loss_v.detach().item() > 1e-20 * loss_initial: for xs, ys in dataloader: optimizer.zero_grad() pred = net(xs) loss_v = torch.norm(pred-ys, p="fro") ** 2 / 2 if epoch == 0: loss_initial = loss_v.detach().item() loss_v.backward() optimizer.step() epoch += 1 # Compute the gradients def loss_fn(input_data, target, model): pred = model(input_data) return torch.norm(pred - target, p="fro") ** 2 / 2 input_data, Y = next(iter(dataloader)) loss_v = loss_fn(input_data, Y, net) grads = torch.autograd.grad(loss_v, net.fc2.weight, create_graph=True) grads = grads[0] # 修正原代码索引错误 # Flatten and concatenate the gradients grad_vector = torch.cat([grad.reshape(-1) for grad in [grads]]) num_params = grad_vector.shape[0] # Compute the Hessian matrix hessian_matrix = torch.zeros((num_params, num_params)).double() for i in range(num_params): optimizer.zero_grad() grad_elem = grad_vector[i].double() # 保留双精度避免精度损失 hessian_row = torch.autograd.grad(grad_elem, net.fc2.weight, retain_graph=True) hessian_row = torch.cat([grad.detach().reshape(-1) for grad in hessian_row]) hessian_matrix[i] = hessian_row # compute condition number of Hessian uh, sh, vh = torch.svd(hessian_matrix) r = torch.linalg.matrix_rank(hessian_matrix, hermitian=True).item() print(f"condition number is {sh[0] / sh[r-1]}_rank of hessian is: {r}") # condition number is 174 # compute Hessian using analytic form w1 = net.fc1.weight.detach().numpy() wx = w1.dot(x) H = np.kron(wx.dot(wx.T), np.eye(output_dim)) U, sh_np, Vh = np.linalg.svd(H, full_matrices=True) r_np = np.linalg.matrix_rank(H) print(f"condition number of by numpy is {sh_np[0] / sh_np[r_np-1]}_rank of hessian is: {r_np}") # condition number is 5345
问题原因分析
代码逻辑错误
- 原代码中梯度索引
grads[1]错误,torch.autograd.grad返回单元素元组,应取grads[0],否则会引发索引越界。 - 将梯度元素转为
float导致双精度张量损失精度,放大数值误差。 - 训练循环未初始化
loss_initial,首次迭代计算相对损失时会出错。 - SVD计算时引用未定义的
torch_hessian变量,实际应使用hessian_matrix。
- 原代码中梯度索引
数值与解析形式的前提差异
- 解析形式是严格全局最小值处的Hessian,但训练后的$W_1$只是接近最优解,并非完全满足全局最优条件,此时Hessian会因$W_1$的偏差产生数值扰动。
- 相对损失$10^{-20}$超出双精度浮点数的有效精度范围(双精度仅15-17位有效数字),此时损失值已为数值噪声,训练无法真正达到该精度,得到的$W_1$并非严格全局最优。
数值计算稳定性问题
- 手动循环计算Hessian每行的方式易累积数值误差,尤其对于500维的大参数矩阵,自动微分的高阶导数计算会因内存和精度限制引入误差。
解决办法
修复代码逻辑错误
- 修正梯度索引、
loss_initial初始化、类型转换、变量引用等问题(见上面修正后的代码)。
- 修正梯度索引、
提升数值计算稳定性
- 使用PyTorch 2.0+的
torch.func.hessian直接计算Hessian,避免手动循环的误差累积:from torch.func import hessian def loss_fn_w2(w2, w1, x, y): pred = w2 @ w1 @ x return torch.norm(pred - y, p="fro") ** 2 / 2 w1_opt = net.fc1.weight.detach() w2_opt = net.fc2.weight.detach() hessian_matrix = hessian(loss_fn_w2, argnums=0)(w2_opt, w1_opt, input_data, Y) hessian_matrix_flat = hessian_matrix.reshape(num_params, num_params)
- 使用PyTorch 2.0+的
调整收敛标准
- 将相对损失收敛标准改为合理范围(如$10^{-12}$),避免超出双精度有效精度导致的数值不稳定。
验证最优解一致性
- 计算全局最优$W_1\star$:当$X$满秩时,$W_1\star$需满足列空间包含$X^\dagger Y\top$($X\dagger$为伪逆),对比训练后$W_1$与理论最优的差异,确认训练是否真正收敛到全局最优。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

