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

使用jax.grad时传入函数参数报错,求微分方程求解解决办法

问题原因

你遇到的TypeError是因为JAX的自动微分函数jax.grad仅支持对数组类参数求导,而你将Python函数dfdx作为loss的参数传入,JAX无法处理这种非JAX兼容的类型。

解决办法

核心思路是避免将Python函数作为jax.grad处理的函数参数,直接在损失函数内部调用目标微分方程的表达式,同时修正代码中的其他潜在问题(如损失计算逻辑、精确解错误等)。

修改后的完整代码
import jax.numpy as jnp
import jax
import matplotlib.pyplot as plt
from tqdm import tqdm
import numpy as np

def softplus(x):
    return jnp.log(1 + jnp.exp(x))

def init_params(key):
    # 修正:将key作为参数传入,避免全局变量依赖
    params = jax.random.normal(key, shape=(241,))
    return params

def linear_model(params, x):
    w0 = params[:80]
    b0 = params[80:160]
    w1 = params[160:240]
    b1 = params[240]
    h = softplus(x*w0 + b0)
    o = jnp.sum(h*w1) + b1
    return o

def dfdx(x, y):
    return -2. * x * y

def loss(initial_condition, params, model, x):
    # 移除derivative参数,直接在内部调用dfdx
    dfdx_model = jax.grad(model, 1)
    dfdx_vect = jax.vmap(dfdx_model, (None, 0))
    model_vect = jax.vmap(model, (None, 0))
    
    y_pred = model_vect(params, x)
    # 用vmap适配dfdx的批量输入(x和y都是数组)
    dfdx_true_vect = jax.vmap(dfdx, (0, 0))
    
    # 修正损失计算:均方误差相加,而非相减
    eq_difference = dfdx_vect(params, x) - dfdx_true_vect(x, y_pred)
    condition_difference = model(params, 0.) - initial_condition
    return jnp.mean(eq_difference ** 2) + jnp.mean(condition_difference ** 2)

key = jax.random.PRNGKey(0)
inputs = np.linspace(0, 1, num=401)
params = init_params(key)

epochs = 2000
learning_rate = 0.0005

# 提前定义梯度函数,指定对params(第1个索引参数)求导
grad_loss = jax.grad(loss, argnums=1)

# 训练循环
for epoch in tqdm(range(epochs)):
    gradient = grad_loss(1., params, linear_model, inputs)
    params -= learning_rate * gradient

# 批量预测
model_vect = jax.vmap(linear_model, (None, 0))
preds = model_vect(params, inputs)

# 修正精确解:原方程y'+2xy=0的解为y=e^(-x²)(初始条件y(0)=1)
plt.plot(inputs, jnp.exp(-inputs**2), label='精确解')
plt.plot(inputs, preds, label='神经网络近似解')
plt.legend()
plt.show()
关键改动说明
  • 移除函数参数:把dfdx从loss的参数列表中移除,直接在函数内部调用,避免JAX处理非兼容类型。
  • 适配批量计算:用jax.vmap包装dfdx,使其能处理批量的x和y_pred输入,保证形状匹配。
  • 修正损失逻辑:将损失从均方误差相减改为相加,确保损失为正,优化方向正确。
  • 明确求导目标:通过jax.grad的argnums=1指定只对模型参数params求导,避免对其他无关参数计算梯度。
  • 修正精确解:原方程的正确解是y=e^(-x²),原代码中符号错误已修正。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 22:42:38