为JAX函数定义非可微参数的通用方法
你提的这个需求太懂行啦——谁没碰到过手动写自定义JVP/VJP写到崩溃的时候?尤其是当函数的微分逻辑复杂到捋不清的时候,能省则省才是正道。我整理了几种不用手动实现微分规则就能标记非可微参数的通用思路,你可以根据自己的场景选:
方法一:利用
jax.grad的argnums参数精准指定求导目标
如果你的场景是直接调用jax.grad求导,其实完全没必要用custom_jvp。直接在jax.grad里通过argnums参数只传入需要求导的参数索引就行,非可微的参数直接排除在外。
举个简单例子:import jax import jax.numpy as jnp def func(a, nondiff_param, b): return a * jnp.sin(nondiff_param) + b # 只对a和b求导,nondiff_param自动被视为非可微 grad_func = jax.grad(func, argnums=(0, 2)) print(grad_func(1.0, jnp.pi/2, 2.0)) # 输出 (1.0, 1.0)这个方法零额外代码,完全靠JAX原生API解决问题,适合大部分简单场景。
方法二:在函数内部用
jax.lax.stop_gradient隔离非可微参数
如果你需要把函数用jax.jit、jax.vmap这类变换包裹,或者不想每次求导都手动指定argnums,可以在函数内部对非可微参数套一层jax.lax.stop_gradient,直接阻断梯度流向这些参数。
示例代码:import jax import jax.numpy as jnp def func(a, nondiff_param, b): # 给非可微参数套上stop_gradient,梯度就传不过来了 safe_nondiff = jax.lax.stop_gradient(nondiff_param) return a * jnp.sin(safe_nondiff) + b grad_func = jax.grad(func, argnums=(0,1,2)) print(grad_func(1.0, jnp.pi/2, 2.0)) # 输出 (1.0, 0.0, 1.0)这里非可微参数的梯度会被置为0,相当于求导时自动忽略它。这个方法灵活性拉满,不管后续用什么JAX变换,都能保证非可微参数的梯度不会被计算。
方法三:用
jax.custom_jvp的简化模式(复用原函数逻辑)
如果你确实需要用custom_jvp来标记nondiff_argnums,其实也不用硬啃复杂的微分推导——JAX允许你只复用原函数的前向逻辑来实现JVP,不用手动写微分公式。
示例代码:from functools import partial import jax import jax.numpy as jnp @partial(jax.custom_jvp, nondiff_argnums=(1,)) def func(a, nondiff_param, b): return a * jnp.sin(nondiff_param) + b # 只需要基于原函数的输入输出计算切向量,不用手动推导微分 @func.defjvp def func_jvp(primals, tangents): a, nondiff_param, b = primals a_tan, _, b_tan = tangents primal_out = func(a, nondiff_param, b) # 直接用原函数里的计算逻辑来推切向量,省掉手动推导的麻烦 tangent_out = jnp.sin(nondiff_param) * a_tan + b_tan return primal_out, tangent_out这个方法虽然还是要写
defjvp,但你不用从头推导微分规则,直接复用原函数里的计算步骤就行,工作量能减少一大半。
总结一下,优先试试前两种方法——jax.grad指定argnums或者函数内部加stop_gradient,这俩都不用碰任何微分规则;如果必须用custom_jvp标记非可微参数,第三种简化模式也能帮你少掉很多头发。
备注:内容来源于stack exchange,提问作者Jingyang Wang

