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

如何统计jax.scipy.optimize.minimize优化中的梯度评估次数?

获取jax.scipy.optimize.minimize梯度评估总次数的实现方法

下面提供三种实用的实现方式,可根据你的优化器类型和需求选择:

方法一:直接读取优化器返回的njev字段

部分梯度类优化器(如L-BFGS-B、BFGS、CG等)会在优化结果中自动记录梯度评估次数,存储在njev字段中,直接读取即可:

import jax
import jax.numpy as jnp
from jax.scipy.optimize import minimize

def objective(x):
    return jnp.sum(x**2)

# 使用L-BFGS-B优化器
result = minimize(objective, jnp.array([1.0, 2.0]), method='L-BFGS-B')

# 打印梯度评估次数
print(f"梯度评估总次数: {result.njev}")

注意:无梯度优化器(如Nelder-Mead)不支持该字段,因为这类方法不需要计算梯度。

方法二:包装梯度函数手动计数

如果优化器不返回njev,或者需要自定义计数逻辑,可以手动包装梯度函数,通过可变对象(如列表)累计计数:

import jax
import jax.numpy as jnp
from jax.scipy.optimize import minimize

# 初始化计数器(用列表存储,利用可变对象特性)
grad_count = [0]

def objective(x):
    return jnp.sum(x**2)

def counted_jac(x):
    # 每次调用梯度函数时计数+1
    grad_count[0] += 1
    return jax.grad(objective)(x)

# 传入自定义的带计数的梯度函数
result = minimize(objective, jnp.array([1.0, 2.0]), method='L-BFGS-B', jac=counted_jac)

print(f"梯度评估总次数: {grad_count[0]}")

方法三:利用jax回调函数计数

通过jax的宿主回调功能,在梯度计算完成后触发计数操作,无需修改原梯度逻辑:

import jax
import jax.numpy as jnp
from jax.scipy.optimize import minimize
from jax.experimental.host_callback import id_tap

grad_count = [0]

# 定义回调计数函数
def increment_count(args, _):
    grad_count[0] += 1

def objective(x):
    return jnp.sum(x**2)

# 包装梯度函数,添加回调
def counted_grad(x):
    grad = jax.grad(objective)(x)
    # 每次计算梯度后触发计数
    id_tap(increment_count, None)
    return grad

result = minimize(objective, jnp.array([1.0, 2.0]), method='L-BFGS-B', jac=counted_grad)

print(f"梯度评估总次数: {grad_count[0]}")

内容的提问来源于stack exchange,提问作者imk

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 13:52:36