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

如何用jax.vmap处理flax TrainStates元组?解决MARL评估报错

问题分析与解决方案

错误原因

你遇到的ValueError是因为jax.vmap要求输入是带有**批量维度(axis 0)**的PyTree结构,但当前的actor_trainstates是两个独立TrainState组成的元组,并非可沿axis 0映射的堆叠结构,vmap无法识别这种异构输入的批量轴。

另外代码里还有两处明显错误需要先修正:

  • PGActor_2的__call__方法中,第一行nn.Dense的输入是未定义的actor_mean,应该替换为输入x;
  • 构建第二个TrainState时,apply_fn误用了network_1.apply,应该对应network_2.apply,否则会因网络结构不匹配导致后续报错。

修正后的完整解决方案

步骤1:修复基础代码错误

先修正语法与逻辑错误,确保代码能正常初始化:

import jax
import jax.numpy as jnp
import flax.linen as nn
from flax.linen.initializers import constant, orthogonal
from flax.training.train_state import TrainState
import optax
import distrax

class PGActor_1(nn.Module):
    @nn.compact
    def __call__(self, x):
        action_dim = 4
        activation = nn.tanh

        actor_mean = nn.Dense(128, kernel_init=orthogonal(jnp.sqrt(2)), bias_init=constant(0.0))(x)
        actor_mean = activation(actor_mean)
        actor_mean = nn.Dense(64, kernel_init=orthogonal(jnp.sqrt(2)), bias_init=constant(0.0))(actor_mean)
        actor_mean = activation(actor_mean)
        actor_mean = nn.Dense(action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(actor_mean)
        return distrax.Categorical(logits=actor_mean)

class PGActor_2(nn.Module):
    @nn.compact
    def __call__(self, x):
        action_dim = 2
        activation = nn.tanh

        # 修正:将未定义的actor_mean替换为输入x
        actor_mean = nn.Dense(64, kernel_init=orthogonal(jnp.sqrt(2)), bias_init=constant(0.0))(x)
        actor_mean = activation(actor_mean)
        actor_mean = nn.Dense(action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(actor_mean)
        return distrax.Categorical(logits=actor_mean)

state= jnp.zeros((1, 5))

network_1 = PGActor_1()
network_1_init_rng = jax.random.PRNGKey(42)
params_1 = network_1.init(network_1_init_rng, state)

network_2 = PGActor_2()
network_2_init_rng = jax.random.PRNGKey(42)
params_2 = network_2.init(network_2_init_rng, state)

tx = optax.chain(
    optax.clip_by_global_norm(1),
    optax.adam(lr=1e-3)
)

# 修正:第二个TrainState使用对应网络的apply_fn
actor_trainstates = (
    TrainState.create(apply_fn=network_1.apply, tx=tx, params=params_1),             
    TrainState.create(apply_fn=network_2.apply, tx=tx, params=params_2)
)

步骤2:方案一 - 堆叠TrainState为带批量轴的PyTree

用jax.tree_map将多个TrainState的对应字段沿axis 0堆叠,形成支持vmap的批量结构:

# 将两个TrainState堆叠成带batch轴的PyTree
batched_trainstate = jax.tree_map(lambda *args: jnp.stack(args), *actor_trainstates)

# 定义批量评估函数
def evaluate_actor(train_state, obs):
    return train_state.apply_fn(train_state.params, obs)

# 应用vmap沿batch轴映射
pis = jax.vmap(evaluate_actor, in_axes=(0, None))(batched_trainstate, state)

# 验证输出
print("Actor 1 logits shape:", pis[0].logits.shape)  # 输出 (1,4)
print("Actor 2 logits shape:", pis[1].logits.shape)  # 输出 (1,2)

步骤3:方案二 - 直接遍历元组映射(适合异构结构)

如果不想堆叠TrainState,可使用jax.lax.map直接遍历元组中的每个TrainState,更适配不同结构的智能体策略:

# 直接对元组中的每个TrainState执行评估
pis = jax.lax.map(lambda x: x.apply_fn(x.params, state), actor_trainstates)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 17:40:09