解决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)。
修改步骤
- 转换参数结构:在主函数中,将原来的列表式集成参数转换为堆叠后的JAX结构:
# 把所有模型的对应参数堆叠,添加模型维度 ensemble_params_stacked = jax.tree_util.tree_map(lambda *args: jnp.stack(args), *ensemble_params) - 使用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
相关产品推荐
相关产品推荐

