如何在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,验证运算一致
方案说明
- 批量层创建:通过
nn.vmap包装nn.Dense,variable_axes={'params':0}让参数增加内部批次轴,split_rngs={'params':True}确保每个内部批次的Dense层用独立随机数初始化,保证参数独立性。 - 无循环调用:
__call__中直接用vmap处理外部批次,调用批量Dense层即可完成运算,完全替代原有循环逻辑。 - 兼容性:初始化时只需传入
do参数即可创建任意数量的内部批次Dense层,符合需求约束。
内容的提问来源于stack exchange,提问作者user137146
相关产品推荐
相关产品推荐

