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

如何使用JAX创建物理信息神经网络(PINN)并实现对输入x的求导?

JAX PINN 模型输入求导的正确实现

首先明确核心问题:

  • jax.grad(params) 报错是因为jax.grad的参数必须是可导函数,而params是模型参数字典,并非函数,自然无法处理。
  • jax.grad(model) 无报错但完全不符合需求:model是Flax Module实例,它的__call__方法第一个参数是self(Module自身),jax.grad(model)实际是尝试对self求导,而JAX无法对Python对象求导,这个操作没有实际意义,也不是你要的对输入x的导数。

正确的求导实现方式

要对模型输出关于输入x求导,需要定义一个明确的前向传播函数,将模型参数和输入x作为输入,输出模型预测值,再针对x执行求导操作。结合你的代码,具体实现如下:

import jax
import jax.numpy as jnp
import flax.linen as fnn
from flax.training import train_state
import optax

class MLP(fnn.Module):
    @fnn.compact
    def __call__(self, x):
        x = fnn.Dense(128)(x)
        x = fnn.relu(x)
        x = fnn.Dense(256)(x)
        x = fnn.relu(x)
        x = fnn.Dense(10)(x)
        return x

model = MLP()
params = model.init(jax.random.PRNGKey(0), jnp.ones([1]))['params']
tx = optax.adam(0.001)
state = train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx)

# 定义接收参数和输入的前向函数
def model_forward(params, x):
    return model.apply({'params': params}, x)

# 1. 对标量输出求导(若模型输出为标量,比如PINN中单个物理量预测)
# argnums=1 指定对函数的第二个参数x求导
grad_x = jax.grad(model_forward, argnums=1)

# 2. 对向量输出求导(你的模型输出维度为10,需计算雅可比矩阵)
# jacfwd是前向模式自动微分,适合输入维度小、输出维度大的场景;jacrev为反向模式,反之
jacobian_x = jax.jacfwd(model_forward, argnums=1)

# 测试求导效果
x_test = jnp.array([1.5])
dx = grad_x(state.params, x_test)
jac = jacobian_x(state.params, x_test)
print(f"输入x的导数(对应标量输出位置):{dx}")
print(f"输出对x的雅可比矩阵:{jac}")

PINN高阶导数计算(如拉普拉斯项)

如果需要计算二阶导数,可嵌套使用jax.grad:

# 计算输出对x的二阶导数
second_grad_x = jax.grad(jax.grad(model_forward, argnums=1), argnums=1)
second_dx = second_grad_x(state.params, x_test)

关键注意点

  • argnums参数:必须明确指定对函数的哪一个输入参数求导,这里model_forward的参数依次是params(第0位)和x(第1位),因此用argnums=1锁定输入x。
  • 向量输出处理:jax.grad仅支持对标量输出求导,若模型输出为多维向量,必须使用jax.jacfwd或jax.jacrev计算雅可比矩阵。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 21:40:21