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

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关闭图追踪后问题仍未解决,我不清楚具体原因,希望了解如何优雅解决该问题。

解决方案

内存占用核心原因

  1. 中间数组维度爆炸:eval_interp1d对300个插值器做vmap,单个x输入会输出(300,)的数组;外层foo再对1e6个x做vmap,会生成(1e6, 300)的中间数组,仅float32类型就占约1.2GB,加上JAX计算过程中的临时缓存,内存占用会进一步飙升。
  2. stop_gradient的局限性:它仅阻止梯度追踪,不会减少正向计算的内存开销——正向过程的中间数组仍会被完整存储。
  3. 冗余维度处理: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 12:27:51