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

神经网络导数计算:如何用矩阵运算替代双重循环实现二阶导数与输入向量的乘法

神经网络导数计算:如何用矩阵运算替代双重循环实现二阶导数与输入向量的乘法

嘿,我仔细看了你的代码和需求,你现在用两层for循环实现的是每个样本、每个输出维度下,输入向量与对应二阶Hessian矩阵的双线性乘积(也就是$result[i,k] = x[i]^T @ second_order_derivative[i,k] @ x[i]$)。完全可以用PyTorch的批量矩阵运算或者Einstein求和来彻底去掉循环,既简洁又能利用硬件并行加速,效率提升会很明显!

先明确核心运算逻辑

你的循环本质上是对每个样本i、每个输出维度k,计算输入向量与对应Hessian矩阵的双线性乘积:

  • x的形状是(2,3):2个样本,每个样本3个输入特征
  • second_order_derivative的形状是(2,3,3,3):2个样本,每个样本对应3个输出的3x3 Hessian矩阵
  • 最终要得到形状为(2,3)的结果:每个样本对应3个输出的计算值

方法1:用Einstein求和(最直观简洁)

PyTorch的torch.einsum可以用类似爱因斯坦求和的符号,直接描述这种多维张量的批量运算,一行代码就能搞定:

# 直接计算批量双线性乘积,结果形状(2,3),和你的循环输出完全一致
result_einsum = torch.einsum('bi,bkij,bj->bk', x, second_order_derivative, x)

对Einstein符号的简单解释:

  • bi:代表x中每个样本b的输入向量维度i(形状(2,3))
  • bkij:代表second_order_derivative中每个样本b、每个输出k对应的Hessian矩阵i→j(形状(2,3,3,3))
  • bj:代表x中每个样本b的输入向量维度j
  • ->bk:最终输出每个样本b、每个输出k的标量结果,形状(2,3)

方法2:用批量矩阵乘法(适合熟悉矩阵运算的场景)

如果更习惯用矩阵乘法的方式,也可以通过调整张量维度,结合torch.bmm(批量矩阵乘法)实现:

# 调整x的维度为(2, 1, 3),适配批量矩阵乘法的输入要求
x_row = x.unsqueeze(1)  # 形状(2,1,3)
# 先计算每个样本每个输出下,x与Hessian的乘积,再和x的列向量相乘得到标量
result_bmm = torch.bmm(
    x_row, 
    torch.bmm(second_order_derivative, x.unsqueeze(2))
).squeeze()

或者拆解步骤更清晰的版本:

# 1. 计算Hessian矩阵与x列向量的批量乘积:(2,3,3,3) @ (2,3,1) → (2,3,3,1)
hessian_x = torch.bmm(second_order_derivative, x.unsqueeze(2))
# 2. 再和x行向量做批量乘积,最后压缩多余维度:(2,1,3) @ (2,3,3,1) → (2,1,1) → 压缩为(2,3)
result_bmm = torch.bmm(x.unsqueeze(1), hessian_x.squeeze(-1)).squeeze()

验证结果一致性

你可以打印对比循环得到的result和两种方法的输出,会发现它们完全相等:

print("循环实现的结果:\n", result)
print("Einstein求和实现的结果:\n", result_einsum)
print("批量矩阵乘法实现的结果:\n", result_bmm)
# 验证数值一致性
print(torch.allclose(result, result_einsum))  # 输出True
print(torch.allclose(result, result_bmm))     # 输出True

额外小优化:二阶导数的计算可以简化

顺便提一句,你当前嵌套两层jacobian计算二阶导数的代码,其实可以用torch.autograd.functional.hessian简化(注意hessian默认对标量输出计算,所以需要针对每个输出维度单独处理):

# 直接计算每个样本每个输出的Hessian矩阵,形状(2,3,3,3),和你的gg结果完全一致
hessian_batch = torch.stack([
    torch.autograd.functional.hessian(lambda xi: calculation(xi)[k], x[i], create_graph=True)
    for i, k in zip(range(x.shape[0]), range(3))
], dim=0)

不过这部分是可选优化,核心还是帮你解决了循环替代的问题。

备注:内容来源于stack exchange,提问作者simon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 15:03:10