如何使用JAX计算批量函数的梯度?含Flax神经网络需求
在JAX中批量计算单变量函数对输入的梯度
你的问题核心在于:jax.grad要求目标函数是标量输出,但你定义的u在输入带批量维度时返回数组输出,导致直接组合vmap和grad时出现维度不兼容问题。下面是两种完全保留批量维度的可行解决方案:
方案1:先定义单样本标量函数,再批量映射
先写一个接受标量输入、返回标量输出的基础函数,再用vmap分别构建批量预测和批量求导的函数:
import jax import jax.numpy as jnp import matplotlib.pyplot as plt # 单样本标量函数:输入标量x,输出标量 u_single = lambda x: jnp.sin(jnp.pi * x) # 批量预测函数:适配批量输入(支持(20,)或(20,1)形状) u_batch = jax.vmap(u_single) # 批量求导函数:对每个样本单独求导 ux_batch = jax.vmap(jax.grad(u_single)) # 生成带批量维度的输入 x = jnp.linspace(-1, 1, 20) # 形状(20,) # x = jnp.expand_dims(jnp.linspace(-1, 1, 20), axis=1) # 形状(20,1)也兼容 plt.plot(x, u_batch(x), label='u(x)') plt.plot(x, ux_batch(x), label="u'(x)") plt.legend() plt.show()
方案2:适配现有批量函数写法
如果不想拆分单样本/批量函数,可以修改u的定义,让它接受单样本数组输入时返回标量,再用vmap处理批量维度:
import jax import jax.numpy as jnp import matplotlib.pyplot as plt x = jnp.expand_dims(jnp.linspace(-1, 1, 20), axis=1) # 保持(20,1)的批量维度 # 调整u:接受单样本数组(如(1,)),返回标量 u = lambda x: jnp.sin(jnp.pi * x).squeeze() # vmap包裹grad(u):对每个(1,)样本单独求导,输出(20,1) ux = jax.vmap(jax.grad(u)) plt.plot(x, u(x), label='u(x)') plt.plot(x, ux(x), label="u'(x)") plt.legend() plt.show()
针对Flax神经网络的适配
如果最终u是Flax实现的神经网络,只需确保网络的单样本前向函数接受标量或单维度数组输入、返回标量输出,再用vmap包裹jax.grad(单样本前向函数),即可直接处理带批量维度的输入:
import flax.linen as nn import jax import jax.numpy as jnp class Net(nn.Module): @nn.compact def __call__(self, x): # x是单样本标量/单维度数组,输出标量 x = nn.Dense(16)(jnp.expand_dims(x, axis=-1)) x = nn.tanh(x) x = nn.Dense(1)(x) return x.squeeze() # 初始化网络 key = jax.random.PRNGKey(42) net = Net() params = net.init(key, jnp.array(0.0)) # 批量预测函数 batch_predict = jax.vmap(lambda p, x: net.apply(p, x), in_axes=(None, 0)) # 批量对输入求导的函数 batch_grad_input = jax.vmap(lambda p, x: jax.grad(lambda x: net.apply(p, x))(x), in_axes=(None, 0)) # 测试批量输入 x_batch = jnp.linspace(-1, 1, 20) predictions = batch_predict(params, x_batch) grads = batch_grad_input(params, x_batch)
内容的提问来源于stack exchange,提问作者al_cc
相关产品推荐
相关产品推荐

