LSTM极简原理 \+ 分步迭代计算实例(自用)
LSTM极简原理 + 分步迭代计算实例(自用)
目录
一、LSTM 极简核心原理
LSTM(长短期记忆网络)是RNN的优化版本,核心解决传统RNN长序列梯度消失、无法记忆长期信息的问题。
其核心本质:用一条长期记忆传送带(细胞状态C)+ 三个门控开关,精准控制记忆的留存、更新、输出,全程只做三件事:
-
遗忘门:扔掉上一步无用的旧记忆
-
输入门:筛选、存入当前步骤的新有效记忆
-
输出门:把当前有效的长期记忆,转化为当前步的输出结果
核心流转逻辑:上一步记忆(h、C)+ 当前输入x → 三门控计算 → 更新新记忆(新C、新h)→ 传入下一步循环计算。
二、前置固定参数与核心通用注意事项(必看)
1. 维度相关核心须知
-
隐藏层维度hidden_dim:这是LSTM的核心超参数,人为设定(常见8/16/32/64/128),维度越大,模型记忆能力越强、参数越多、越容易过拟合。本文为方便手算,统一设置hidden_dim=1,工业场景需根据数据量上调维度。
-
输入特征维度feature_dim:单时间步的输入特征数(单特征=1,多特征=2/3/多),本文设定feature_dim=1。
-
拼接规则(核心):每一步计算的输入,永远是 上一步隐状态h(t-1) + 当前输入x(t) 拼接,拼接后维度 = hidden_dim + feature_dim,所有权重矩阵的维度严格匹配该拼接维度。
2. 训练与迭代轮次须知
-
序列轮次(时间步seq_len):LSTM没有固定轮次,由任务场景决定。时序预测常见seq_len=10/20/30(用前10/20/30步数据预测下一步),文本任务由句子单词数决定。
-
训练迭代轮次epoch:常规训练轮次为20~100轮,小数据集30轮左右即可收敛,大数据集可适当增加,避免欠拟合。
3. 初始化须知
- 模型首次计算(t0)无前置记忆,需手动初始化h0、c0,工业场景默认初始化为全0向量,本文为直观计算,设定固定非0初始值。
4. 固定公式与激活函数
基础计算流程:拼接向量→四门控线性计算→激活→更新细胞状态→输出隐状态
激活函数:sigmoid(x)=1/(1+e⁻ˣ)、tanh(x)=(eˣ-e⁻ˣ)/(eˣ+e⁻ˣ)
固定权重偏置(全程不变,保证迭代连贯性):Wf=Wi=Wc=Wo=[0.2, 0.4],bf=bi=bc=bo=0.1
三、完整两步迭代计算案例(连贯循环)
初始设定
初始记忆状态:h0=0.1,c0=0.2
时序输入序列(2个时间步):第一步输入x1=0.5,第二步输入x2=0.8
第一步:t1 时刻,输入 x1=0.5,输入来源:h0、c0、x1
1、向量拼接:[h0, x1] = [0.1, 0.5](维度匹配:1+1=2,与权重维度适配)
2、遗忘门计算:控制保留多少上一步细胞记忆c0
Wf·拼接向量 + bf = 0.2×0.1 + 0.4×0.5 + 0.1 = 0.32,f1=σ(0.32)≈0.5794
3、输入门计算:控制存入多少当前新记忆
Wi·拼接向量 + bi = 0.32,i1=σ(0.32)≈0.5794
4、候选记忆计算:当前时刻可写入的新记忆内容
Wc·拼接向量 + bc = 0.32,ṽc1=tanh(0.32)≈0.3090
5、更新长期细胞状态(核心记忆更新):c1 = f1×c0 + i1×ṽc1
c1 = 0.5794×0.2 + 0.5794×0.3090 = 0.1159 + 0.1790 = 0.2949
6、输出门计算:控制多少长期记忆转为当前输出
Wo·拼接向量 + bo = 0.32,o1=σ(0.32)≈0.5794
7、更新当前隐状态(短期输出,传入下一步):h1 = o1×tanh(c1)
tanh(0.2949)≈0.2864,h1=0.5794×0.2864≈0.1660
t1输出结果(传入下一步):h1=0.1660,c1=0.2949
第二步:t2 时刻,输入 x2=0.8,输入来源:h1、c1、x2(承接上一步结果)
1、向量拼接:[h1, x2] = [0.1660, 0.8](延续拼接规则:上一步h + 当前x)
2、遗忘门计算
Wf·拼接向量 + bf = 0.2×0.1660 + 0.4×0.8 + 0.1 = 0.4532,f2=σ(0.4532)≈0.6114
3、输入门计算
Wi·拼接向量 + bi = 0.4532,i2=σ(0.4532)≈0.6114
4、候选记忆计算
Wc·拼接向量 + bc = 0.4532,ṽc2=tanh(0.4532)≈0.4247
5、更新长期细胞状态:c2 = f2×c1 + i2×ṽc2
c2 = 0.6114×0.2949 + 0.6114×0.4247 = 0.1803 + 0.2597 = 0.4400
6、输出门计算
Wo·拼接向量 + bo = 0.4532,o2=σ(0.4532)≈0.6114
7、更新当前隐状态:h2 = o2×tanh(c2)
tanh(0.4400)≈0.4140,h2=0.6114×0.4140≈0.2531
t2输出结果:h2=0.2531,c2=0.4400
四、迭代循环核心逻辑总结
1. 每一个时间步的计算,完全依赖上一步的h、c结果,实现时序记忆传递;
2. 所有门控共享同一份拼接向量,权重偏置全程固定,仅输入和记忆状态动态更新;
3. 细胞状态C负责长期记忆留存,几乎无梯度损耗,是LSTM适配长序列的核心;隐状态h负责输出当前步结果,参与下一步迭代。
五、补充实操关键注意事项
1. 维度升级规则:手算用hidden_dim=1,实际使用时若设置hidden_dim=32,则h、c均为32维向量,权重矩阵同步升级为 [32, feature_dim+32],拼接逻辑不变;
2. 输入归一化:LSTM依赖sigmoid、tanh激活,输入数据必须归一化到[0,1]或[-1,1],否则激活函数饱和、模型无法学习;
3. 序列长度适配:seq_len不宜过长(超过50步建议分层LSTM),否则仍会出现轻微记忆衰减;
4. 权重更新机制:本文为前向推理计算,模型训练时会根据误差反向更新W、b参数,多轮迭代后权重收敛,实现精准预测。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)