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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 01:24:05