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无法拆分,未成功
技术问询
ScipyBoundedMinimize首次迭代前/期间速度慢的潜在原因是什么?- 针对大规模参数、大数据量且含插值的复杂模型场景,
jax中是否存在更快的替代优化算法? - 我是否误解了
optax.adam的并行化方式?该场景下有哪些可行的并行化策略? - 提供的代码片段中是否存在可优化点(如向量化)以提升性能?
补充信息
- 硬件: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
相关产品推荐
相关产品推荐

