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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 10:30:50