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

如何在Flax模型中用vmap映射多个Dense实例?避免循环

实现批次专属参数的Flax模型优化方案

原始实现代码

from jax import random,vmap
from jax import numpy as jnp
import pprint
import flax.linen as nn  # 补充缺失的导入

def f(s,layers,do,dx):
    x = jnp.zeros((do,dx))
    for i,layer in enumerate(layers):
        x=x.at[i].set( layer( s[i] ) )
    return x

class net(nn.Module):
    dx: int 
    do: int 
    def setup(self):
        self.layers = [ nn.Dense( self.dx, use_bias=False )
                        for _ in range(self.do) ]
    def __call__(self, s):
        x = vmap(f,in_axes=(0,None,None,None))(s,self.layers,self.do,self.dx)
        return x

if __name__ == '__main__':
    seed = 123
    key = random.PRNGKey( seed )
    key,subkey = random.split( key )
    outer_batches = 4
    s_observations = 5 # AKA the inner batch
    x_features = 2
    s_features = 3
    s_shape = (outer_batches,s_observations, s_features)
    s = random.uniform( subkey, s_shape )

    key,subkey = random.split( key )    
    model = net(x_features,s_observations)
    p = model.init( subkey, s )
    x = model.apply( p, s )    

    params = p['params']
    pkernels = jnp.array([params[key]['kernel'] for key in params.keys()])
    x_=jnp.zeros((outer_batches,s_observations,x_features))
    
    g = vmap(vmap(lambda a,b: a@b),in_axes=(0,None))
    
    x_=g(s,pkernels)
    print('s shape:',s.shape)
    print('p shape:',pkernels.shape)
    print('x shape:',x.shape)
    print('x_ shape:',x_.shape)
    print('sum of difference:',jnp.sum(x-x_))

问题描述

需要在模型中设置批次专属参数:存在长度为do的“内部批次”,每个批次元素对应一个flax.linen.Dense实例;外部批次仅负责向这些层传入多组数据实例。

当前通过在setup方法中创建Dense实例列表实现需求,在__call__中遍历列表填充数组,再用jax.vmap封装循环。同时已写出等价的矩阵乘法逻辑(函数g)明确运算目标。

希望用jax.vmap调用直接替代__call__中的循环,但直接向vmap传入列表、或将多个Dense实例放入JAX数组都会报错。需要找到替代列表的方式存储这些Dense实例,且初始化模型时需支持创建任意数量的Dense实例。

优化方案:用Flax的vmap批量创建Dense层

可以通过flax.linen.vmap批量生成Dense层,参数会被自动组织成批次化结构,无需手动维护列表,同时天然支持jax.vmap的批量调用。

优化后的代码

from jax import random, vmap
from jax import numpy as jnp
import flax.linen as nn

class net(nn.Module):
    dx: int 
    do: int 
    def setup(self):
        # 用flax.linen.vmap批量创建do个Dense层,参数维度增加内部批次轴
        self.batch_dense = nn.vmap(
            nn.Dense,
            in_axes=None,  # 输入无内部批次轴(后续由外部vmap处理)
            out_axes=0,    # 输出增加内部批次轴
            variable_axes={'params': 0},  # 参数增加内部批次轴,实现每个内部批次独立参数
            split_rngs={'params': True}   # 每个Dense层用独立随机数初始化
        )(features=self.dx, use_bias=False)
    
    def __call__(self, s):
        # s形状:(outer_batches, do, s_features)
        # 对外部批次做vmap,直接调用批量Dense层完成运算
        return vmap(self.batch_dense)(s)

if __name__ == '__main__':
    seed = 123
    key = random.PRNGKey(seed)
    key, subkey = random.split(key)
    
    outer_batches = 4
    do = 5  # 内部批次数量
    dx = 2  # 输出特征数
    s_features = 3  # 输入特征数
    
    s_shape = (outer_batches, do, s_features)
    s = random.uniform(subkey, s_shape)
    
    key, subkey = random.split(key)
    model = net(dx=dx, do=do)
    params = model.init(subkey, s)
    x = model.apply(params, s)
    
    # 提取批次化核参数,验证运算等价性
    batch_kernel = params['params']['batch_dense']['kernel']
    x_ = vmap(lambda s_batch: s_batch @ batch_kernel)(s)
    
    print('s shape:', s.shape)
    print('batch_kernel shape:', batch_kernel.shape)
    print('x shape:', x.shape)
    print('x_ shape:', x_.shape)
    print('sum of difference:', jnp.sum(x - x_))  # 输出应为0,验证运算一致

方案说明

  1. 批量层创建:通过nn.vmap包装nn.Dense,variable_axes={'params':0}让参数增加内部批次轴,split_rngs={'params':True}确保每个内部批次的Dense层用独立随机数初始化,保证参数独立性。
  2. 无循环调用:__call__中直接用vmap处理外部批次,调用批量Dense层即可完成运算,完全替代原有循环逻辑。
  3. 兼容性:初始化时只需传入do参数即可创建任意数量的内部批次Dense层,符合需求约束。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 05:53:17