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

网络输出对输入的梯度计算代码错误排查请求

代码问题排查:网络输入输出梯度计算错误

问题背景

这段代码目标是计算网络输出相对于输入的梯度,但返回值不符合预期。要求矩阵B必须为如下结构的3×3矩阵:

[0          0           0
 0    d²c/dx²    d²(cu)/dxdy
 0  d²c/dxdy    d²(cu)/dy²]

原代码

import jax.numpy as jnp
from jax import jacfwd, jacrev

class NCL(object):
    def __init__(self, network, mass_constant=2):
        self.network = network
        self.mc = mass_constant

    def A(self, x, params):
            u_v = self.network(x,params[0])[:-1]
            c, cu = self.network(x, params[0])[:2]  # Extract c and cu
            dc_dx = jacfwd(lambda x: c)(x)
            print(dc_dx)
            d2c_dx2 = jacfwd(lambda x: dc_dx)(x)[1]
            d2c_dxdy = jacfwd(lambda x: dc_dx)(x)[2]
            dcu_dy = jacfwd(lambda x: cu)(x)[2]
            d2cu_dy2 = jacfwd(lambda x: dcu_dy)(x)[2]
            d2cu_dydx = jacfwd(lambda x: dcu_dy)(x)[1]
            
            I_2 = d2c_dx2 + d2cu_dydx
            I_3 = d2c_dxdy + d2cu_dy2


            I = [[0.0],[I_2],[I_3]]
            print("d",dc_dx)
            N = len(x)
            B = jnp.zeros((N,N))
            diag_idx = jnp.diag_indices(N,1)
            B = B.at[diag_idx].set(I)

           
            A = jnp.zeros((N,N))
            idx = jnp.triu_indices(N,1)
            A = A.at[idx].set(u_v)
            return A - A.T + B

# Example usage
network = lambda x, params: jnp.array([x[0]*x[1]*x[1] + params[0], x[1]*x[1]*x[1] * params[1], x[2] * params[2],x[2]])

ncl_instance = NCL(network)

# Example input
x_input = jnp.array([1.0, 2.0, 3.0])

# Example parameters
params_input = jnp.array([0.1, 0.2, 0.3])

# Test the A function
result_A = ncl_instance.A(x_input, [params_input])

print("Result A:\n", result_A)

错误点与修正方案

1. 二阶导数计算逻辑错误

原代码对一阶导数结果再次求导并直接索引,逻辑完全错误。正确做法是通过两次jacfwd直接获取函数的Hessian矩阵,再提取对应位置的二阶导数:

  • 对于c(x, y),d²c/dx²是Hessian矩阵的(0,0)位置,d²c/dxdy是(0,1)位置;
  • 对于cu(x, y),d²(cu)/dxdy是Hessian矩阵的(0,1)位置,d²(cu)/dy²是(1,1)位置。

2. 矩阵B构造完全不符合要求

原代码错误使用偏移对角索引,实际需要手动构造符合要求的3×3矩阵:

  • 第一行全为0;
  • 第二行设置[0, d²c/dx², d²(cu)/dxdy];
  • 第三行设置[0, d²c/dxdy, d²(cu)/dy²]。

3. A矩阵构造维度不匹配

原代码中u_v是长度为3的向量,但jnp.triu_indices(N,1)返回的是2个索引对((0,1),(0,2)),无法直接赋值。需将u_v的元素对应到上三角的(0,1),(0,2),(1,2)位置,再转为反对称矩阵。

4. 冗余的网络调用

原代码两次调用self.network,合并为一次调用可提升效率。

修正后的代码

import jax.numpy as jnp
from jax import jacfwd, jacrev

class NCL(object):
    def __init__(self, network, mass_constant=2):
        self.network = network
        self.mc = mass_constant

    def A(self, x, params):
        # 单次调用网络获取所有输出
        net_output = self.network(x, params[0])
        u_v = net_output[:-1]
        c, cu = net_output[:2]

        # 计算c的Hessian矩阵(二阶导数)
        hess_c = jacfwd(jacfwd(lambda x: self.network(x, params[0])[0]))(x)
        d2c_dx2 = hess_c[0, 0]
        d2c_dxdy = hess_c[0, 1]

        # 计算cu的Hessian矩阵(二阶导数)
        hess_cu = jacfwd(jacfwd(lambda x: self.network(x, params[0])[1]))(x)
        d2cu_dxdy = hess_cu[0, 1]
        d2cu_dy2 = hess_cu[1, 1]

        # 构造矩阵B
        B = jnp.zeros((3, 3))
        B = B.at[1, 1].set(d2c_dx2)
        B = B.at[1, 2].set(d2cu_dxdy)
        B = B.at[2, 1].set(d2c_dxdy)
        B = B.at[2, 2].set(d2cu_dy2)

        # 构造反对称矩阵A
        A = jnp.zeros((3, 3))
        # 将u_v元素对应到上三角指定位置
        A = A.at[0, 1].set(u_v[0])
        A = A.at[0, 2].set(u_v[1])
        A = A.at[1, 2].set(u_v[2])
        # 转为反对称矩阵
        A = A - A.T

        return A + B

# Example usage
network = lambda x, params: jnp.array([x[0]*x[1]*x[1] + params[0], x[1]*x[1]*x[1] * params[1], x[2] * params[2], x[2]])

ncl_instance = NCL(network)

# Example input
x_input = jnp.array([1.0, 2.0, 3.0])

# Example parameters
params_input = jnp.array([0.1, 0.2, 0.3])

# Test the A function
result_A = ncl_instance.A(x_input, [params_input])

print("Result A:\n", result_A)

说明

修正后的代码:

  • 正确计算了所需的二阶导数;
  • 构造出符合要求的矩阵B;
  • 修复了A矩阵的维度不匹配问题,生成正确的反对称矩阵;
  • 减少了冗余计算,提升了代码效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 15:01:27