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

如何在Flux.jl中使用LSTM?LSTM(in::Integer, out::Integer)参数含义是什么?

Flux.jl LSTM使用说明

参数含义

  • in:输入特征的维度,即单步输入张量的第一个维度长度。比如你从视频每帧提取到的特征向量长度是256,这里就填256。
  • out:LSTM层输出隐藏状态的维度,你可以根据任务需求自由调整,值越大LSTM的记忆容量越高,对应的计算开销也越大。

基础使用方法

Flux的LSTM是有状态层,会自动保留之前输入的上下文记忆,不需要手动传递隐藏状态和细胞状态。

1. 层定义

导入Flux后直接调用构造函数即可,示例:

using Flux
# 输入特征维度128,输出隐藏状态维度64
lstm_layer = LSTM(128, 64)

2. 输入输出格式

LSTM默认接收的单步输入维度为(特征维度, 批次大小),如果是单样本输入,批次大小可以填1。
单步调用示例:

# 单样本单帧输入,特征维度128,批次大小1
single_frame = rand(Float32, 128, 1)
# 输出隐藏状态维度为(64, 1)
hidden_state = lstm_layer(single_frame)

处理序列输入时,直接按时间步依次传入输入即可,LSTM会自动更新内部状态:

# 共16帧的视频序列,每帧特征维度128,批次大小4
frame_sequence = [rand(Float32, 128, 4) for _ in 1:16]
# 依次传入所有帧,得到每一步的隐藏状态
hidden_sequence = [lstm_layer(frame) for frame in frame_sequence]

3. 状态重置

处理完独立的序列或批次后,必须调用Flux.reset!清空LSTM的内部状态,否则上一个序列的记忆会残留到下一次计算,导致结果错误:

Flux.reset!(lstm_layer)

4. 完整模型示例(视频分类场景)

using Flux

# 定义端到端模型:LSTM提取时序特征 + 全连接层分类
video_clf = Chain(
    LSTM(128, 64), # 输入每帧特征128维,输出64维隐藏状态
    Dense(64, 10), # 接全连接层输出10分类结果
    softmax
)

# 批次输入:共8个视频,每个视频20帧
batch_seq = [rand(Float32, 128, 8) for _ in 1:20]
# 前向传播
for frame in batch_seq
    result = video_clf(frame)
end
# 最终result维度为(10, 8),即每个样本对应10类的概率
# 重置模型状态,准备下一批次计算
Flux.reset!(video_clf)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 20:42:03