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

使用pmap多GPU运行含循环的高开销函数时出现OOM错误

解决JAX pmap多GPU运行时循环导致的OOM问题

问题背景

我有一个高开销函数expensive_func,需要对形状为(N, m)的数组inputs中的N组输入执行计算。计划用4个GPU(ngpu=4,可整除N)分片处理,每个GPU负责N/4个案例,最后汇总结果。

单GPU串行运行大批次时无内存问题,仅耗时较长;但用pmap实现多GPU并行时,即便每个GPU处理量未增加,仍出现内存不足(OOM)错误。怀疑pmap编译时将循环展开,导致内存需求飙升。已尝试分片输入避免冗余拷贝,但希望保留单GPU循环的低内存特性——无pmap时循环运行正常,加pmap就OOM,仿佛循环被并行展开了。

问题原因

JAX的即时编译(JIT)机制包括pmap,会将Python中的for循环编译时展开,把循环转化为一次性计算所有迭代的向量化操作。这意味着原本单GPU串行循环中,每次迭代只保留当前计算的中间结果,现在所有迭代的中间结果会同时驻留在内存中,直接导致内存占用暴涨,触发OOM。

解决方案

要维持循环的串行内存占用特性,需要使用JAX原生的动态循环操作,比如jax.lax.fori_loop或jax.lax.scan,它们会在运行时串行执行循环迭代,不会在编译时展开,从而避免内存占用激增。

方案1:使用jax.lax.fori_loop

适合简单的索引式循环,直接指定迭代范围和迭代逻辑:

import jax
import jax.numpy as jnp
from jax import pmap

N = 8
m = 10
inputs = jnp.array(jnp.arange(N*m).reshape(N, m), dtype=jnp.float32)

def expensive_func(inp):
    return jnp.sum(inp ** 2)

# 用jax.lax.fori_loop替代Python循环,实现串行累加
def singledevice_func(local_inputs):
    def body(i, accum):
        val = expensive_func(local_inputs[i])
        return accum + val
    # 从0到batch_size-1迭代,初始accum为0.0
    return jax.lax.fori_loop(0, local_inputs.shape[0], body, 0.0)

pmapped = pmap(singledevice_func, in_axes=0)

inputs_sharded = inputs.reshape(jax.device_count(), -1, m)
accum_dev = pmapped(inputs_sharded)
accum_final = jnp.sum(accum_dev)

print(accum_final)

方案2:使用jax.lax.scan

更灵活,适合处理序列输入的场景,自动遍历输入数组的第一个维度:

import jax
import jax.numpy as jnp
from jax import pmap

N = 8
m = 10
inputs = jnp.array(jnp.arange(N*m).reshape(N, m), dtype=jnp.float32)

def expensive_func(inp):
    return jnp.sum(inp ** 2)

# 用jax.lax.scan实现串行累加
def singledevice_func(local_inputs):
    def step(accum, inp):
        val = expensive_func(inp)
        return accum + val, None  # scan需要返回(新状态, 输出),这里输出用None忽略
    accum_final, _ = jax.lax.scan(step, 0.0, local_inputs)
    return accum_final

pmapped = pmap(singledevice_func, in_axes=0)

inputs_sharded = inputs.reshape(jax.device_count(), -1, m)
accum_dev = pmapped(inputs_sharded)
accum_final = jnp.sum(accum_dev)

print(accum_final)

补充说明

  • jax.lax.fori_loop和jax.lax.scan都会让JAX在运行时串行执行循环,每次迭代只保留当前的累加值,内存占用和单GPU串行循环一致;
  • 同时结合pmap实现多GPU的分片并行,既保留低内存特性,又能利用多GPU加速计算。

内容的提问来源于stack exchange,提问作者evening silver fox

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.11 16:05:06