PyTorch中计算简单神经网络参数Hessian的报错问题咨询
问题解决:PyTorch计算神经网络权重Hessian报错
TypeError: forward() missing 1 required positional argument: 'target' 错误根因
你调用torch.autograd.functional.hessian的传参逻辑有问题:
- 该接口的第一个入参要求是可调用函数,且函数的输入恰好是你要求二阶导的变量,也就是这里的网络权重
- 你直接传入了
nn.CrossEntropyLoss的实例,首先损失函数本身需要接收模型输出、真实标签两个入参,你只传了扁平化的权重参数,自然会缺参数报错;其次损失函数和网络权重没有绑定,接口无法知道权重的变化如何影响损失的计算结果。
解决方法
你需要把「用当前权重计算模型输出→计算损失」的整个逻辑封装成一个仅接收扁平化权重作为输入的函数,固定计算损失用到的样本和标签,具体实现如下:
首先提前记录参数的形状和长度,方便后续把一维的权重张量恢复为模型各层的参数格式:
# 固定计算Hessian用到的输入样本和对应标签,可按需替换为全数据集或指定批次 fixed_input = x_var.detach().clone() fixed_target = y_var.detach().clone() # 记录各层参数的形状、元素数量 param_shapes = [p.shape for p in net.parameters()] param_numels = [p.numel() for p in net.parameters()]
然后封装损失计算函数:
def compute_loss(params_flat): # 将传入的一维权重张量拆分,赋值给模型对应层的参数 ptr = 0 for param, shape, numel in zip(net.parameters(), param_shapes, param_numels): param.data = params_flat[ptr:ptr+numel].view(shape) ptr += numel # 前向传播计算损失 model_output = net(fixed_input) return loss_func(model_output, fixed_target)
最后调用接口计算Hessian即可:
hessian = torch.autograd.functional.hessian(compute_loss, param_list, create_graph=True)
注意事项
- Hessian矩阵的维度是「参数总数量 × 参数总数量」,如果你的模型参数量较大,很容易出现显存溢出,建议先把隐藏层尺寸设为极小值(比如2、5)验证逻辑正确性
- 如果你不需要对Hessian做后续的求导操作,可将
create_graph设为False,大幅降低显存占用和计算耗时 - 上述代码计算的是
fixed_input、fixed_target对应样本集的Hessian,若需要计算全数据集的平均Hessian,只需把封装函数内的损失替换为全数据集的平均损失即可
内容的提问来源于stack exchange,提问作者Kasra
相关产品推荐
相关产品推荐

