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

如何用JAX/Flax的nnx.vmap实现带独立参数的nnx.Module集合向量化

Flax NNX实现带独立参数的神经网络集合(等价于Linen的nn.vmap(variable_axes={"params":0}))

背景

我已基于Flax Linen实现可正常运行的向量化Q网络,集合中每个critic拥有独立参数,输出沿第一维度(n_critics)堆叠。现在尝试迁移至Flax NNX,但NNX的nnx.vmapAPI与语义和Linen差异较大,无法确定如何正确创建带独立参数的Module集合。

Flax Linen实现代码

import jax.numpy as jnp
from flax import linen as nn
from typing import Optional

class QNetwork(nn.Module):
    action_dim: int
    rew_dim: int
    hidden_dim: int = 256

    @nn.compact
    def __call__(self, obs, w, deterministic: bool):
        x = nn.Dense(self.hidden_dim)(obs)
        x = nn.relu(x)
        x = nn.Dense(self.action_dim * self.rew_dim)(x)
        return x


class VectorQNetwork(nn.Module):
    action_dim: int
    rew_dim: int
    n_critics: int = 2

    @nn.compact
    def __call__(self, obs: jnp.ndarray, w: jnp.ndarray, deterministic: bool):
        vmap_critic = nn.vmap(
            QNetwork,
            variable_axes={"params": 0},        # 每个critic使用独立参数
            split_rngs={"params": True},        # 初始化使用不同随机数流
            in_axes=None,
            out_axes=0,
            axis_size=self.n_critics,
        )

        q_values = vmap_critic(
            action_dim=self.action_dim,
            rew_dim=self.rew_dim,
        )(obs, w, deterministic)

        return q_values.reshape(
            (self.n_critics, -1, self.action_dim, self.rew_dim)
        )

迁移NNX遇到的问题

找不到等价方式完成:

  • 创建nnx.Module集合
  • 确保每个critic拥有独立参数
  • 保留相同输出堆叠行为

NNX的nnx.vmap未像Linen那样暴露variable_axes、split_rngs或axis_size参数,集合式用法示例稀缺。我自己实现了可运行的NNX代码,但速度比Linen版本慢2倍以上,并非最优方案:

我当前的NNX实现(性能不佳)

class VectorQNetwork(nnx.Module):
    """Vectorized QNetwork."""

    def __init__(self,
                 obs_dim: int,
                 action_dim: int,
                 rew_dim: int,
                 use_layer_norm: bool = True,
                 dropout_rate: Optional[float] = 0.01,
                 n_critics: int = 2,
                 num_hidden_layers: int = 4,
                 hidden_dim: int = 256,
                 image_obs: bool = False,
                 rngs: nnx.Rngs = None):
        self.action_dim = action_dim
        self.rew_dim = rew_dim
        self.use_layer_norm = use_layer_norm
        self.dropout_rate = dropout_rate
        self.n_critics = n_critics
        self.num_hidden_layers = num_hidden_layers
        self.hidden_dim = hidden_dim
        self.image_obs = image_obs

        @nnx.split_rngs(splits=n_critics)
        @nnx.vmap(axis_size=n_critics)
        def create_q_net(rngs):
            return QNetwork(
                obs_dim=obs_dim,
                action_dim=action_dim,
                rew_dim=rew_dim,
                dropout_rate=dropout_rate,
                use_layer_norm=use_layer_norm,
                num_hidden_layers=num_hidden_layers,
                hidden_dim=hidden_dim,
                image_obs=image_obs,
                rngs=rngs,
            )
        
        self.q_net = create_q_net(rngs)

    def __call__(self, obs: jnp.ndarray, w: jnp.ndarray):
        """Forward pass of the Q network."""
        def apply_q(q_net):
            return q_net(obs, w)
        
        q_values = nnx.vmap(apply_q)(self.q_net)
        return q_values.reshape((self.n_critics, -1, self.action_dim, self.rew_dim))

正确的NNX实现方式

要等价于Linen中nn.vmap(..., variable_axes={"params":0})的效果,核心是利用nnx.vmap对Module的参数维度进行映射,同时通过nnx.split_rngs确保每个critic初始化使用独立随机数流,且在forward时直接对Module的forward方法做vmap,避免额外函数包装开销。

优化后的NNX代码

import jax.numpy as jnp
import nnx
from typing import Optional

class QNetwork(nnx.Module):
    def __init__(self,
                 obs_dim: int,
                 action_dim: int,
                 rew_dim: int,
                 hidden_dim: int = 256,
                 use_layer_norm: bool = True,
                 dropout_rate: Optional[float] = 0.01,
                 num_hidden_layers: int = 4,
                 image_obs: bool = False,
                 rngs: nnx.Rngs = None):
        super().__init__()
        self.action_dim = action_dim
        self.rew_dim = rew_dim
        self.hidden_dim = hidden_dim
        self.use_layer_norm = use_layer_norm
        self.dropout_rate = dropout_rate
        self.num_hidden_layers = num_hidden_layers
        self.image_obs = image_obs

        # 构建网络层,与原Linen实现逻辑对齐
        self.layers = []
        # 输入层
        self.layers.append(nnx.Dense(hidden_dim, in_features=obs_dim, rngs=rngs))
        if use_layer_norm:
            self.layers.append(nnx.LayerNorm(rngs=rngs))
        self.layers.append(nnx.relu)
        # 隐藏层
        for _ in range(num_hidden_layers - 1):
            self.layers.append(nnx.Dense(hidden_dim, rngs=rngs))
            if use_layer_norm:
                self.layers.append(nnx.LayerNorm(rngs=rngs))
            self.layers.append(nnx.relu)
            if dropout_rate > 0:
                self.layers.append(nnx.Dropout(rate=dropout_rate, rngs=rngs))
        # 输出层
        self.layers.append(nnx.Dense(action_dim * rew_dim, rngs=rngs))

    def __call__(self, obs: jnp.ndarray, w: jnp.ndarray, deterministic: bool = True):
        x = obs
        for layer in self.layers:
            if isinstance(layer, nnx.Dropout):
                x = layer(x, deterministic=deterministic)
            else:
                x = layer(x)
        return x

class VectorQNetwork(nnx.Module):
    def __init__(self,
                 obs_dim: int,
                 action_dim: int,
                 rew_dim: int,
                 n_critics: int = 2,
                 hidden_dim: int = 256,
                 use_layer_norm: bool = True,
                 dropout_rate: Optional[float] = 0.01,
                 num_hidden_layers: int = 4,
                 image_obs: bool = False,
                 rngs: nnx.Rngs = None):
        super().__init__()
        self.n_critics = n_critics
        self.action_dim = action_dim
        self.rew_dim = rew_dim

        # 核心实现:创建带独立参数的QNetwork集合
        # 1. split_rngs拆分随机数流,确保每个critic初始化独立
        # 2. vmap将QNetwork的参数映射到第0维度,实现参数独立
        @nnx.split_rngs(splits=n_critics)
        @nnx.vmap(axis_size=n_critics)
        def create_vectorized_qnet(rngs):
            return QNetwork(
                obs_dim=obs_dim,
                action_dim=action_dim,
                rew_dim=rew_dim,
                hidden_dim=hidden_dim,
                use_layer_norm=use_layer_norm,
                dropout_rate=dropout_rate,
                num_hidden_layers=num_hidden_layers,
                image_obs=image_obs,
                rngs=rngs,
            )

        self.q_networks = create_vectorized_qnet(rngs)

    def __call__(self, obs: jnp.ndarray, w: jnp.ndarray, deterministic: bool = True):
        # 直接对QNetwork的forward方法做vmap,减少额外开销
        q_values = nnx.vmap(lambda net: net(obs, w, deterministic))(self.q_networks)
        # 保持与Linen版本一致的输出形状
        return q_values.reshape((self.n_critics, -1, self.action_dim, self.rew_dim))

关键优化点说明

  1. 参数独立化:@nnx.vmap(axis_size=n_critics)装饰QNetwork创建函数,自动将所有参数(如Dense的kernel、bias)映射到第0维度(长度为n_critics),等价于Linen的variable_axes={"params":0}。
  2. 独立初始化:@nnx.split_rngs(splits=n_critics)将输入rng拆分为n_critics个独立流,每个QNetwork用不同rng初始化,等价于Linen的split_rngs={"params":True}。
  3. 性能提升:forward时直接对QNetwork的__call__方法做vmap,避免额外函数包装,减少JAX调度开销,解决速度慢的问题。
  4. 输出一致性:最终reshape操作与Linen版本完全一致,输出沿n_critics维度堆叠。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.11 20:13:13