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

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}")

原示例代码的问题说明

  1. 硬编码随机键:原代码中直接使用jax.random.key(0)初始化carry,不符合Flax的rng管理规范,会导致初始化流程出错。
  2. 缺少序列长度指定:部分2023年Flax版本中,nn.scan包装的LSTMCell需要显式指定length参数,否则会引发维度推断错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 07:07:03