Jax是否支持对向量化参数的指定索引位置求导?
JAX对向量指定索引求导的解决方案
JAX原生的grad接口的argnums参数仅支持指定整个输入参数作为求导对象,不支持直接通过嵌套元组的形式指定某个参数内部的索引位置求导,你之前尝试的argnums=((0,),)写法不属于argnums的合法入参格式,因此无法生效。
你可以通过以下两种常用方案实现需求:
- 方案1:先对整个向量参数求导,再提取对应索引的梯度值,适合绝大多数场景,写法最简单
import jax.numpy as jnp from jax import grad def test_func(a): return a[0]**a[1] a = jnp.array([2.0, 3.0]) # 先得到整个向量的梯度,形状和a一致 full_grad = grad(test_func)(a) # 提取对应索引的梯度,比如对a[0]的梯度就是full_grad[0] grad_a0 = full_grad[0] print(grad_a0) # 输出为12.0,对应3*2^(3-1)的计算结果
- 方案2:封装外层函数,将需要求导的索引位置单独拆为独立参数,适合需要固定对某个索引求导、后续多次调用的场景
import jax.numpy as jnp from jax import grad def test_func(a): return a[0]**a[1] # 封装函数,把要单独求导的a[0]提为第一个参数 wrapped_test = lambda a0, a1: test_func(jnp.array([a0, a1])) # 对第一个参数(也就是原a的0号索引)求导 grad_fn = grad(wrapped_test, argnums=0) a = jnp.array([2.0, 3.0]) print(grad_fn(a[0], a[1])) # 输出同样为12.0
内容的提问来源于stack exchange,提问作者user654123
相关产品推荐
相关产品推荐

