如何使用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
相关产品推荐
相关产品推荐

