You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

问题原因分析

  1. 代码逻辑错误

    • 原代码中梯度索引grads[1]错误,torch.autograd.grad返回单元素元组,应取grads[0],否则会引发索引越界。
    • 将梯度元素转为float导致双精度张量损失精度,放大数值误差。
    • 训练循环未初始化loss_initial,首次迭代计算相对损失时会出错。
    • SVD计算时引用未定义的torch_hessian变量,实际应使用hessian_matrix。
  2. 数值与解析形式的前提差异

    • 解析形式是严格全局最小值处的Hessian,但训练后的$W_1$只是接近最优解,并非完全满足全局最优条件,此时Hessian会因$W_1$的偏差产生数值扰动。
    • 相对损失$10^{-20}$超出双精度浮点数的有效精度范围(双精度仅15-17位有效数字),此时损失值已为数值噪声,训练无法真正达到该精度,得到的$W_1$并非严格全局最优。
  3. 数值计算稳定性问题

    • 手动循环计算Hessian每行的方式易累积数值误差,尤其对于500维的大参数矩阵,自动微分的高阶导数计算会因内存和精度限制引入误差。

解决办法

  1. 修复代码逻辑错误

    • 修正梯度索引、loss_initial初始化、类型转换、变量引用等问题(见上面修正后的代码)。
  2. 提升数值计算稳定性

    • 使用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)
      
  3. 调整收敛标准

    • 将相对损失收敛标准改为合理范围(如$10^{-12}$),避免超出双精度有效精度导致的数值不稳定。
  4. 验证最优解一致性

    • 计算全局最优$W_1\star$:当$X$满秩时,$W_1\star$需满足列空间包含$X^\dagger Y\top$($X\dagger$为伪逆),对比训练后$W_1$与理论最优的差异,确认训练是否真正收敛到全局最优。

内容的提问来源于stack exchange,提问作者Simon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 02:47:33