无需调用model.apply(),如何访问Flax类/模块的子模块?
解决Flax中Autoencoder子模块单独调用的问题
你可以通过在autoencoder类中添加单独调用子模块的方法来实现需求,既能通过一次model.init获取所有参数,又能单独调用每个子模块,具体操作如下:
修改后的Autoencoder类
class autoencoder(nn.Module): hidden_dim: int z_dim: int output_dim: int def setup(self): self.encoder = encoder(self.hidden_dim, self.z_dim) self.decoder = decoder(self.hidden_dim, self.output_dim) # 其他4个子模块也在这里完成初始化 def __call__(self, x): z = self.encoder(x) y = self.decoder(z) return y # 添加单独调用encoder的方法 def encode(self, x): return self.encoder(x) # 添加单独调用decoder的方法 def decode(self, z): return self.decoder(z) # 其他子模块对应的调用方法可以按同样模式添加
调用方式
- 正常执行完整前向传播:
params = autoencoder(hidden_dim=256, z_dim=64, output_dim=784).init(rng, x_sample) recon = autoencoder().apply(params, x_input) - 单独调用encoder:
z = autoencoder().apply(params, x_input, method=autoencoder.encode) - 单独调用decoder:
recon_z = autoencoder().apply(params, z_sample, method=autoencoder.decode)
这种方式下,所有子模块的参数会统一存放在同一个params字典中(结构对应子模块名称,比如params['encoder']、params['decoder']),不需要分开初始化,同时能灵活单独调用任意子模块。
虽然Flax文档的Future Work提到了更直接的子模块访问方式,但当前通过添加方法的方式已经能很好满足你的需求,不管是6个还是更多子模块,都可以用同样的方法扩展。
内容的提问来源于stack exchange,提问作者Jabby
相关产品推荐
相关产品推荐

