网络输出对输入的梯度计算代码错误排查请求
代码问题排查:网络输入输出梯度计算错误
问题背景
这段代码目标是计算网络输出相对于输入的梯度,但返回值不符合预期。要求矩阵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
相关产品推荐
相关产品推荐

