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

如何用JAX/Flax.nnx的vmap实现多模型推理并行化?

问题:使用JAX vmap并行处理多个NNX模型时的报错解决

原始代码与需求

现有如下预测函数,通过循环遍历模型列表处理对应输入:

from flax import nnx
from jax import Array
from typing import List

def predict(models: List[nnx.Module], imgs: Array):
    for i, model in enumerate(models):
       prediction = model(imgs[i, ...])

希望用jax.vmap或nnx.vmap实现GPU全并行处理,避免循环,但直接传入模型列表会报错。

尝试的方案与报错

尝试用普通jax.vmap包装单模型预测函数:

from jax import vmap

def predict_single(model, img):
   return model(img)

predict = vmap(predict_single)

触发报错:

ValueError: vmap was requested to map its argument along axis 0, which implies that its rank should be at least 1, but is only 0 (its shape is ())

解决方法

普通jax.vmap无法直接处理nnx.Module列表,因为nnx.Module是PyTree结构,列表形式的模型集合没有被标记为带批量维度的结构。需要使用NNX专门提供的nnx.vmap来处理模型的批量映射:

方法1:用nnx.vmap直接包装预测函数

from flax import nnx
from jax import Array
from typing import List

def predict_single(model: nnx.Module, img: Array) -> Array:
    return model(img)

# 使用nnx.vmap,指定对model和img的第0维进行映射
predict = nnx.vmap(predict_single, in_axes=(0, 0))

# 调用示例:models是List[nnx.Module],imgs是形状为(N, ...)的Array(N与模型数量一致)
predictions = predict(models, imgs)

方法2:先将模型列表转为批量模型再处理

如果需要更直观的批量模型调用,可以先把模型列表转换成带批量维度的NNX模型,再直接处理输入:

from flax import nnx
from jax import Array
from typing import List

# 将模型列表转为批量模型
batched_model = nnx.vmap(lambda model: model, in_axes=0)(models)

# 直接用批量模型处理批量输入(imgs形状为(N, ...))
predictions = batched_model(imgs)

原理说明

nnx.vmap是专门针对NNX模块设计的批量映射工具,它能自动识别并处理nnx.Module的PyTree结构,将列表形式的模型集合视为第0维的批量数据,从而实现与输入图片的并行映射,避免了普通jax.vmap无法识别模型批量维度的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 18:33:12