JAX中无法计算含梯度的Lambda函数梯度问题求助
问题原因
你代码里的核心问题是:f_x和f_xx是提前计算好的f[0]的梯度函数,属于静态引用。当定义f_next时直接调用这些预存的梯度函数,JAX在追踪f_next的计算图时,无法将它们关联到f_old的动态导数计算逻辑,导致后续求f_next的梯度时计算链断裂。
解决方案
把f_next内部的预存梯度替换为动态计算的梯度,让JAX能完整追踪整个依赖链:
import jax import jax.numpy as jnp # Model parameters γ = 1.5 k = 0.1 μY = 0.03 σ = 0.03 λ = 0.1 ωb = μY/λ # PDE params. σω = σ dt =0.01 IC = lambda ω: jnp.exp(-(1-γ)*ω) f = [IC] f_old = f[0] # 动态计算f_old的一阶、二阶梯度,而非提前预存 f_next = lambda ω: f_old(ω) + 100*dt * ( (0.5*σω**2)*jax.grad(jax.grad(f_old))(ω) - λ*(ω-ωb)*jax.grad(f_old)(ω) - k*f_old(ω) + jnp.exp(-(1-γ)*ω)) print(f_next(0.)) f.append(f_next) f_x= jax.grad(f[1]) # 计算f_next的一阶导数 print(f_x(0.))
关键修改说明
- 移除了提前定义的
f_x和f_xx变量,改为在f_next内部通过jax.grad实时计算f_old的梯度 - 这样JAX在构建
f_next的计算图时,能明确识别出梯度与f_old的依赖关系,后续求f_next的梯度时就能正常追踪完整的计算流程
内容的提问来源于stack exchange,提问作者Marco
相关产品推荐
相关产品推荐

