如何结合nnx.split_rngs与jax.tree.map创建不同输出维度的Linear层?
使用nnx.split_rngs结合jax.tree.map生成多输出维度Linear层
问题描述
我正在学习使用nnx.split_rngs,想找一段结合nnx.split_rngs与jax.tree.map的代码,用来生成任意数量、不同out_features的Linear层。
现有代码通过包装函数my_linear_wrapper实现了将2维输入映射到2至6维输出的Linear层树结构,但希望改用类似@nnx.split_rngs装饰器的方式实现,想问能不能在my_linear上用nnx.split_rngs,为nnx.Linear的rng参数实现映射?原代码如下:
import jax from flax import nnx from functools import partial if __name__ == '__main__': session_sizes = { 'a':2, 'b':3, 'c':4, 'd':5, 'e':6, } dz = 2 rngs = nnx.Rngs(0) my_linear = partial( nnx.Linear, use_bias = False, in_features = dz, rngs=rngs ) def my_linear_wrapper(a): return my_linear( out_features=a ) q_s = jax.tree.map(my_linear_wrapper, session_sizes) for k in session_sizes.keys(): print(q_s[k].kernel)
解决方案
完全可以用nnx.split_rngs实现这个需求,核心是通过装饰器为每个Linear层分配独立的RNG分支,配合jax.tree.map遍历生成不同输出维度的层。修改后的代码如下:
import jax from flax import nnx if __name__ == '__main__': session_sizes = { 'a': 2, 'b': 3, 'c': 4, 'd': 5, 'e': 6, } dz = 2 rngs = nnx.Rngs(0) # 用nnx.split_rngs装饰器包装Linear创建逻辑 @nnx.split_rngs def create_linear(out_features: int, rngs: nnx.Rngs): return nnx.Linear( in_features=dz, out_features=out_features, use_bias=False, rngs=rngs ) # 直接用jax.tree.map遍历session_sizes生成Linear层树结构 q_s = jax.tree.map(create_linear, session_sizes) for k in session_sizes.keys(): print(f"Layer {k} kernel shape: {q_s[k].kernel.shape}") print(q_s[k].kernel)
关键说明
@nnx.split_rngs装饰器会自动为每次函数调用拆分出独立的RNG子流,确保每个Linear层的权重初始化使用不同的随机种子,避免权重重复。- 我们将Linear层的创建逻辑直接封装在
create_linear函数中,不再需要额外的包装函数my_linear_wrapper,代码更简洁。 jax.tree.map会遍历session_sizes字典的每个值,传入create_linear生成对应输出维度的Linear层,最终得到和原代码结构一致的层树。
内容的提问来源于stack exchange,提问作者jworrell
相关产品推荐
相关产品推荐

