JAX嵌套vmap内存占用过高问题排查与优化咨询
问题
我正在使用JAX处理三维网格上大量插值器的计算任务,遵循JAX标准实践,先编写单批次输入代码,最后通过嵌套vmap遍历所有插值器和评估网格点。下方是简化后的示例代码:
from collections import namedtuple from functools import partial import jax.numpy as jnp from jax import vmap from jax.lax import dynamic_slice, stop_gradient interpolation_params = namedtuple("interpolation_params", ["a", "dx", "f", "lb", "ub"]) @partial(vmap, in_axes=(None, None, 0)) def init_1d_interpolation_params(a, dx, f): f = jnp.pad(f, 1) lb, ub = a, a + (f.shape[0] - 1) * dx return interpolation_params(a=a, dx=dx, f=f, lb=lb, ub=ub) @partial(vmap, in_axes=(None, 0)) def eval_interp1d(x, interpolation_params): A = jnp.array([-1.0 / 16, 9.0 / 16, 9.0 / 16, -1.0 / 16]) B = jnp.array([1.0 / 24, -9.0 / 8, 9.0 / 8, -1.0 / 24]) C = jnp.array([1.0 / 4, -1.0 / 4, -1.0 / 4, 1.0 / 4]) D = jnp.array([-1.0 / 6, 1.0 / 2, -1.0 / 2, 1.0 / 6]) x = ( jnp.minimum(jnp.maximum(x, interpolation_params.lb), interpolation_params.ub) - interpolation_params.a ) ix = jnp.atleast_1d(jnp.array(x // interpolation_params.dx, int)) ratx = x / interpolation_params.dx - (ix + 0.5) asx = A + ratx * (B + ratx * (C + ratx * D)) return jnp.dot(dynamic_slice(interpolation_params.f, ix, (4,)), asx) # Init 300 interpolants on a uniform grid with 4096 points x = jnp.linspace(0, 1, 4096) f = x**2 ff = jnp.repeat(f.reshape(1, -1), 300, axis=0) params = init_1d_interpolation_params(x[0], x[1] - x[0], ff) @partial(vmap, in_axes=(0, None)) def foo(x, interpolation_params): g_x = (eval_interp1d(x, interpolation_params)) ** 2 return jnp.sum(g_x) large_x_array = stop_gradient(jnp.repeat(jnp.array([0.0]), 100**3)) foo(large_x_array, params)
运行该代码时出现了高达14GB的内存占用,这令我困惑。起初我认为是JAX自动微分后端的计算图追踪导致的,其规模理论上与params和large_x_array的笛卡尔积相当,但使用stop_gradient关闭图追踪后问题仍未解决,我不清楚具体原因,希望了解如何优雅解决该问题。
解决方案
内存占用核心原因
- 中间数组维度爆炸:
eval_interp1d对300个插值器做vmap,单个x输入会输出(300,)的数组;外层foo再对1e6个x做vmap,会生成(1e6, 300)的中间数组,仅float32类型就占约1.2GB,加上JAX计算过程中的临时缓存,内存占用会进一步飙升。 stop_gradient的局限性:它仅阻止梯度追踪,不会减少正向计算的内存开销——正向过程的中间数组仍会被完整存储。- 冗余维度处理:
ix = jnp.atleast_1d(...)给标量索引添加了不必要的维度,增加了dynamic_slice的内存开销。
优化方案
1. 调整vmap嵌套逻辑,避免大数组累积
将求和操作嵌入内层vmap,每个x的计算仅生成(300,)的数组,求和后立即释放内存,不会累积(1e6,300)的超大数组:
# 先定义单x单插值器的计算逻辑 def eval_interp1d_single(x, params): A = jnp.array([-1.0 / 16, 9.0 / 16, 9.0 / 16, -1.0 / 16]) B = jnp.array([1.0 / 24, -9.0 / 8, 9.0 / 8, -1.0 / 24]) C = jnp.array([1.0 / 4, -1.0 / 4, -1.0 / 4, 1.0 / 4]) D = jnp.array([-1.0 / 6, 1.0 / 2, -1.0 / 2, 1.0 / 6]) x_clamped = jnp.minimum(jnp.maximum(x, params.lb), params.ub) - params.a ix = jnp.array(x_clamped // params.dx, dtype=int) ratx = x_clamped / params.dx - (ix + 0.5) asx = A + ratx * (B + ratx * (C + ratx * D)) # 修正dynamic_slice的索引格式,去掉不必要的atleast_1d return jnp.dot(dynamic_slice(params.f, (ix,), (4,)), asx) # 单x对所有插值器计算并求和 def compute_single_x_sum(x, params): interp_vals = vmap(eval_interp1d_single, in_axes=(None, 0))(x, params) return jnp.sum(interp_vals ** 2) # 外层vmap遍历所有x,最后总求和 foo_optimized = vmap(compute_single_x_sum, in_axes=(0, None)) total_sum = jnp.sum(foo_optimized(large_x_array, params))
2. 启用JIT编译,自动复用内存
给计算函数加上@jax.jit装饰器,JAX会编译出高效的循环逻辑,自动复用中间缓冲区,大幅降低内存占用:
from jax import jit @jit @partial(vmap, in_axes=(0, None)) def foo(x, interpolation_params): g_x = (eval_interp1d(x, interpolation_params)) ** 2 return jnp.sum(g_x)
3. 手动分批次处理x数组
如果1e6个x点的规模仍超出内存承载,可手动分批次计算,逐步累加结果:
batch_size = 10000 num_batches = len(large_x_array) // batch_size total_sum = 0.0 # 分批次计算 for i in range(num_batches): x_batch = large_x_array[i*batch_size : (i+1)*batch_size] batch_sum = jnp.sum(foo(x_batch, params)) total_sum += batch_sum # 处理剩余数据 if len(large_x_array) % batch_size != 0: x_batch = large_x_array[num_batches*batch_size:] total_sum += jnp.sum(foo(x_batch, params))
内容的提问来源于stack exchange,提问作者crypty
相关产品推荐
相关产品推荐

