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

解决JAX vmap并行化集成网络损失计算的轴尺寸不一致错误

解决JAX vmap并行计算集成网络损失时的轴尺寸不一致错误

错误原因

你遇到的ValueError根源在于:ensemble_params是Python列表嵌套的参数结构,JAX的vmap无法自动识别这是模型维度的批量结构。当你指定in_axes=(0, None, None)时,vmap会误将各层参数本身的维度(比如第一层权重的3、第二层权重的4)当作要映射的轴,导致轴尺寸不匹配的错误。

解决方案

需要将列表形式的集成参数转换为带模型批量轴的JAX数组树——用jax.tree_util.tree_map配合jnp.stack,把所有模型的对应参数沿轴0堆叠,给整个参数结构添加一个模型维度(大小为num_models)。

修改步骤

  1. 转换参数结构:在主函数中,将原来的列表式集成参数转换为堆叠后的JAX结构:
    # 把所有模型的对应参数堆叠,添加模型维度
    ensemble_params_stacked = jax.tree_util.tree_map(lambda *args: jnp.stack(args), *ensemble_params)
    
  2. 使用vmap计算损失:直接用堆叠后的参数调用vmap包装的损失函数即可:
    ensemble_loss = jax.vmap(mse_loss, in_axes=(0, None, None))
    losses = ensemble_loss(ensemble_params_stacked, x, y)
    

完整修改后的代码

import jax
from jax import Array
from jax import random
import jax.numpy as jnp
from jax.tree_util import tree_map

def layer_params(dim_in: int, dim_out: int, key: Array) -> tuple[Array]:
    w_key, b_key = random.split(key=key)
    weights = random.normal(key=w_key, shape=(dim_out, dim_in))
    biases = random.normal(key=w_key, shape=(dim_out,))
    return weights, biases

def init_params(layer_dims: list[int], key: Array) -> list[tuple[Array]]:
    keys = random.split(key=key, num=len(layer_dims))
    params = []
    for dim_in, dim_out, key in zip(layer_dims[:-1], layer_dims[1:], keys):
        params.append(layer_params(dim_in=dim_in, dim_out=dim_out, key=key))
    return params

def init_ensemble(key: Array, num_models: int, layer_dims: list[int]) -> list:
    keys = random.split(key=key, num=num_models)
    models = [init_params(layer_dims=layer_dims, key=key) for key in keys]
    return models

def relu(x):
  return jnp.maximum(0, x)

def predict(params, image):
  activations = image
  for w, b in params[:-1]:
    outputs = jnp.dot(w, activations) + b
    activations = relu(outputs)
  final_w, final_b = params[-1]
  logits = jnp.dot(final_w, activations) + final_b
  return logits

batched_predict = jax.vmap(predict, in_axes=(None, 0))

def mse_loss(params, inputs, targets):
    preds = batched_predict(params, inputs)
    loss = jnp.mean((targets - preds) ** 2)
    return loss

if __name__ == "__main__":

    num_models = 4
    dim_in = 2
    dim_out = 4
    layer_dims = [dim_in, 3, dim_out]
    batch_size = 2

    key = random.PRNGKey(seed=1)
    key, subkey = random.split(key)
    ensemble_params = init_ensemble(key=subkey, num_models=num_models, layer_dims=layer_dims)

    key_x, key_y = random.split(key)
    x = random.normal(key=key_x, shape=(batch_size, dim_in))
    y = random.normal(key=key_y, shape=(batch_size, dim_out))

    # 原循环计算(用于验证)
    for params in ensemble_params:
        loss = mse_loss(params, inputs=x, targets=y)
        print(f"{loss = }")

    # 转换参数结构并并行计算
    ensemble_params_stacked = tree_map(lambda *args: jnp.stack(args), *ensemble_params)
    ensemble_loss = jax.vmap(mse_loss, in_axes=(0, None, None))
    losses = ensemble_loss(ensemble_params_stacked, x, y)
    print(f"{losses = }")  # 结果与循环计算完全一致

原理说明

jax.tree_util.tree_map会递归遍历参数的嵌套结构(每个模型的参数是[ (w1,b1), (w2,b2) ]的形式),对所有模型的对应参数(比如所有模型的w1)执行jnp.stack,将它们沿轴0堆叠成形状为(num_models, dim_out, dim_in)的数组。这样整个参数结构的每个叶子节点都带有模型维度,vmap就能正确沿着这个维度并行计算每个模型的损失,不会再混淆参数本身的维度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 07:27:42