2023年FLAX LSTM使用问题:官方示例报错求可用实现
2023年Flax LSTM时间序列输入可用示例
方案一:使用Flax封装好的nn.LSTM层(推荐)
Flax在2023年的版本中已经提供了开箱即用的nn.LSTM层,无需手动用nn.scan包装LSTMCell,代码更简洁且符合最新API规范:
import flax.linen as nn import jax import jax.numpy as jnp class SimpleLSTMModel(nn.Module): hidden_dim: int output_dim: int = None # 可选:如果需要输出层则指定 @nn.compact def __call__(self, x): # 输入形状:(batch_size, seq_len, input_dim) lstm = nn.LSTM(self.hidden_dim) # 使用Flax内置rng系统初始化carry状态,避免硬编码随机键 carry = lstm.initialize_carry(self.make_rng('lstm'), x.shape[:1]) carry, outputs = lstm(carry, x) # 可选:添加全连接输出层 if self.output_dim is not None: outputs = nn.Dense(self.output_dim)(outputs) return outputs # 测试代码 batch_size = 4 seq_len = 12 input_dim = 7 hidden_dim = 32 x = jnp.ones((batch_size, seq_len, input_dim)) model = SimpleLSTMModel(hidden_dim=hidden_dim) # 初始化模型参数,传入所需的rng键字典 key = jax.random.key(0) y, params = model.init_with_output({'params': key, 'lstm': key}, x) print(f"输入形状: {x.shape}") print(f"输出形状: {y.shape}")
方案二:手动用nn.scan包装LSTMCell(兼容底层需求)
如果需要手动控制scan逻辑,以下是适配2023年Flax API的修正版本:
import flax.linen as nn import jax import jax.numpy as jnp class ScanLSTMModel(nn.Module): hidden_dim: int @nn.compact def __call__(self, x): # 配置scan包装LSTMCell,明确指定序列长度 ScanLSTM = nn.scan( nn.LSTMCell, variable_broadcast="params", split_rngs={"params": False}, in_axes=1, # 序列维度位于axis=1 out_axes=1, length=x.shape[1] # 显式指定序列长度,避免自动推断错误 ) lstm_cell = ScanLSTM(self.hidden_dim) # 用Flax rng系统初始化carry,输入形状对应batch维度 carry = lstm_cell.initialize_carry(self.make_rng('lstm'), x.shape[:1]) carry, outputs = lstm_cell(carry, x) return outputs # 测试代码 x = jnp.ones((4, 12, 7)) model = ScanLSTMModel(hidden_dim=32) key = jax.random.key(0) y, params = model.init_with_output({'params': key, 'lstm': key}, x) print(f"输入形状: {x.shape}") print(f"输出形状: {y.shape}")
原示例代码的问题说明
- 硬编码随机键:原代码中直接使用
jax.random.key(0)初始化carry,不符合Flax的rng管理规范,会导致初始化流程出错。 - 缺少序列长度指定:部分2023年Flax版本中,
nn.scan包装的LSTMCell需要显式指定length参数,否则会引发维度推断错误。
内容的提问来源于stack exchange,提问作者al_cc
相关产品推荐
相关产品推荐

