如何用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
相关产品推荐
相关产品推荐

