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

JAX大规模模型优化提速咨询:ScipyBoundedMinimize性能瓶颈排查

JAX模型优化性能优化问题

我正在使用jax框架优化一个包含插值操作的复杂模型,该模型需要拟合4800个观测数据点。当前采用jaxopt.ScipyBoundedMinimize进行优化,100次迭代耗时约30秒,且大部分时间消耗在首次迭代期间或开始前。必要数据(idc, sg and cpcs)可通过压缩包获取。

import jax.numpy as jnp
import time as ela_time
from jaxopt import ScipyBoundedMinimize
import optax
import jax
import pickle


file1 = open('idc.pkl', 'rb')
idc = pickle.load(file1)
file1.close()

file2 = open('sg.pkl', 'rb')
sg = pickle.load(file2)
file2.close()

file3 = open('cpcs.pkl', 'rb')
cpcs = pickle.load(file3)
file3.close()


def model(fssc, fssh, time, rv, amp):

    fssp = 1.0 - (fssc + fssh)

    ivis = cpcs['common'][time]['ivis']
    areas = cpcs['common'][time]['areas']
    mus = cpcs['common'][time]['mus']

    vels = idc['vels'].copy()

    ldfs_phot = cpcs['line'][time]['ldfs_phot']
    ldfs_cool = cpcs['line'][time]['ldfs_cool']
    ldfs_hot = cpcs['line'][time]['ldfs_hot']

    lps_phot = cpcs['line'][time]['lps_phot']
    lps_cool = cpcs['line'][time]['lps_cool']
    lps_hot = cpcs['line'][time]['lps_hot']

    lis_phot = cpcs['line'][time]['lis_phot']
    lis_cool = cpcs['line'][time]['lis_cool']
    lis_hot = cpcs['line'][time]['lis_hot']

    coeffs_phot = lis_phot * ldfs_phot * areas * mus
    wgt_phot = coeffs_phot * fssp[ivis]
    wgtn_phot = jnp.sum(wgt_phot)

    coeffs_cool = lis_cool * ldfs_cool * areas * mus
    wgt_cool = coeffs_cool * fssc[ivis]
    wgtn_cool = jnp.sum(wgt_cool)

    coeffs_hot = lis_hot * ldfs_hot * areas * mus
    wgt_hot = coeffs_hot * fssh[ivis]
    wgtn_hot = jnp.sum(wgt_hot)

    prf = jnp.sum(wgt_phot[:, None] * lps_phot + wgt_cool[:, None] * lps_cool + wgt_hot[:, None] * lps_hot, axis=0)
    prf /= wgtn_phot + wgtn_cool + wgtn_hot

    prf = jnp.interp(vels, vels + rv, prf)

    prf = prf + amp

    avg = jnp.mean(prf)

    prf = prf / avg

    return prf


def loss(x0s, lmbd):

    noes = sg['noes']

    noo = len(idc['times'])

    fssc = x0s[:noes]
    fssh = x0s[noes: 2 * noes]
    fssp = 1.0 - (fssc + fssh)
    rv = x0s[2 * noes: 2 * noes + noo]
    amp = x0s[2 * noes + noo: 2 * noes + 2 * noo]

    chisq = 0
    for i, itime in enumerate(idc['times']):
        oprf = idc['data'][itime]['prf']
        oprf_errs = idc['data'][itime]['errs']

        nop = len(oprf)

        sprf = model(fssc=fssc, fssh=fssh, time=itime, rv=rv[i], amp=amp[i])

        chisq += jnp.sum(((oprf - sprf) / oprf_errs) ** 2) / (noo * nop)

    wp = sg['grid_areas'] / jnp.max(sg['grid_areas'])

    mem = jnp.sum(wp * (fssc * jnp.log(fssc / 1e-5) + fssh * jnp.log(fssh / 1e-5) +
                    (1.0 - fssp) * jnp.log((1.0 - fssp) / (1.0 - 1e-5)))) / sg['noes']

    ftot = chisq + lmbd * mem

    return ftot


if __name__ == '__main__':

    # idc: a dictionary containing observational data (150 x 32)
    # sg and cpcs: dictionaries with related coefficients

    noes = sg['noes']
    lmbd = 1.0
    maxiter = 1000
    tol = 1e-5

    fss = jnp.ones(2 * noes) * 1e-5
    x0s = jnp.hstack((fss, jnp.zeros(len(idc['times']) * 2)))

    minx0s = [1e-5] * (2 * noes) + [-jnp.inf] * len(idc['times']) * 2
    maxx0s = [1.0 - 1e-5] * (2 * noes) + [jnp.inf] * len(idc['times']) * 2

    bounds = (minx0s, maxx0s)

    start = ela_time.time()

    optimizer = ScipyBoundedMinimize(fun=loss, maxiter=maxiter, tol=tol, method='L-BFGS-B',
                                 options={'disp': True})
    x0s, info = optimizer.run(x0s, bounds,  lmbd)

    # optimizer = optax.adam(learning_rate=0.1)
    # optimizer_state = optimizer.init(x0s)
    #
    # for i in range(1, maxiter + 1):
    #
    #     print('ITERATION -->', i)
    #
    #     gradients = jax.grad(loss)(x0s, lmbd)
    #     updates, optimizer_state = optimizer.update(gradients, optimizer_state, x0s)
    #     x0s = optax.apply_updates(x0s, updates)
    #     x0s = jnp.clip(x0s, jnp.array(minx0s), jnp.array(maxx0s))
    #     print('Objective function: {:.3E}'.format(loss(x0s, lmbd)))

    end = ela_time.time()

    print(end - start)   # total elapsed time: ~30 seconds

相关信息说明

  • 自由参数(x0s)数量:5263
  • 数据:观测数据存储在idc字典中(共4800个数据点)
  • 模型:定义在model函数中,包含插值操作
  • 已尝试的优化方法:
    • jaxopt.ScipyBoundedMinimize(L-BFGS-B方法):速度较慢,约30秒,大部分时间消耗在首次迭代期间或开始前
    • optax.adam:速度过慢,约200秒
  • 并行化尝试:尝试对optax.adam进行并行化,但由于模型固有特性,x0s无法拆分,未成功

技术问询

  1. ScipyBoundedMinimize首次迭代前/期间速度慢的潜在原因是什么?
  2. 针对大规模参数、大数据量且含插值的复杂模型场景,jax中是否存在更快的替代优化算法?
  3. 我是否误解了optax.adam的并行化方式?该场景下有哪些可行的并行化策略?
  4. 提供的代码片段中是否存在可优化点(如向量化)以提升性能?

补充信息

  • 硬件:Intel® Core™ i7-9750H CPU @ 2.60GHz × 12,16 GiB RAM(笔记本电脑)
  • 软件:操作系统Ubuntu 22.04,Python 3.10.12,JAX 0.4.25,optax 0.2.1

问题解答

1. ScipyBoundedMinimize首次迭代慢的原因

  • JAX的JIT编译开销:首次调用loss或其梯度时,JAX需要将Python函数编译成XLA可执行代码,这个过程会占用大量时间。L-BFGS-B在首次迭代前会计算初始损失和梯度,触发完整的JIT编译,这是最主要的耗时来源。
  • 参数规模大:5263个参数的梯度计算本身复杂度高,首次编译需要处理所有相关的运算图,生成优化的机器码,耗时自然更长。
  • 数据格式转换开销:加载的pickle数据是numpy格式,首次迭代时JAX需要将其转换为JAX数组格式,涉及内存拷贝和格式适配,会增加额外耗时。

2. JAX中更快的替代优化算法

  • jaxopt.LBFGS:JAX原生实现的L-BFGS,避免了Python与Scipy的跨进程交互开销,能更好地利用JAX的编译优化,大规模参数场景下效率优于Scipy封装版本。
  • jaxopt.Adam:JAXopt封装的Adam优化器,内置边界约束处理(通过prox参数),相比手动用optax实现,梯度计算和更新逻辑更高效。
  • trust-constr方法:通过jaxopt.ScipyMinimize选择该方法,适合大规模约束优化,收敛速度比L-BFGS-B更快,尤其在目标函数非凸性较强的场景。
  • 小批量随机优化:如果损失函数支持小批量近似,用optax.adamw配合小批量数据训练,单次迭代计算量大幅降低,总耗时可能低于全批量优化。

3. optax.adam的并行化误解与可行策略

你误解了并行化的方向:参数无法拆分不代表不能并行化计算。JAX的并行化核心是运算的向量化和设备并行,而非参数拆分。可行策略:

  • 向量化时间步循环:用jax.vmap替换loss中遍历idc['times']的Python循环,让JAX自动并行处理所有时间步的损失计算,充分利用CPU多核心。
  • 启用多线程并行:确保JAX开启CPU多线程(默认已开启,可通过jax.config.update('jax_cpu_multi_thread', True)确认),让XLA自动将运算分配到多个CPU核心执行。
  • GPU加速:如果有GPU可用,将数据和参数转移到GPU,JAX会自动将运算并行到CUDA核心,计算速度可提升数倍。

4. 代码中的可优化点

  • 向量化loss中的时间步循环:
    def single_time_loss(itime, rv_i, amp_i):
        oprf = idc['data'][itime]['prf']
        oprf_errs = idc['data'][itime]['errs']
        sprf = model(fssc, fssh, itime, rv_i, amp_i)
        return jnp.sum(((oprf - sprf) / oprf_errs) ** 2)
    
    times = jnp.array(list(idc['times']))
    rv_arr = jnp.array(rv)
    amp_arr = jnp.array(amp)
    chisq_per_time = jax.vmap(single_time_loss)(times, rv_arr, amp_arr)
    chisq = jnp.mean(chisq_per_time) / len(oprf)
    
    用jax.vmap替代Python循环,消除解释型循环的开销,同时实现并行计算。
  • 提前转换数据为JAX数组:加载pickle后,将所有numpy数组转换为JAX数组(jnp.array()),避免函数内部重复转换:
    idc['vels'] = jnp.array(idc['vels'])
    for time in idc['times']:
        idc['data'][time]['prf'] = jnp.array(idc['data'][time]['prf'])
        idc['data'][time]['errs'] = jnp.array(idc['data'][time]['errs'])
    # 对sg、cpcs做类似转换
    
  • JIT编译损失函数:给loss添加@jax.jit装饰器,提前编译运算图,后续迭代直接复用编译后的代码:
    @jax.jit(static_argnums=(1,))
    def loss(x0s, lmbd):
        # 原有代码
    
  • 移除不必要的复制:vels = idc['vels'].copy()改为vels = idc['vels'],JAX数组不可变,复制无意义且浪费内存。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 22:00:55