2025-11-15 学习记录--Python-LSTM模型定义(PyTorch)

LSTM模型定义(PyTorch)

  • LSTM(Long Short-Term Memory)长短期记忆网络
    是 RNN(循环神经网络)的一种改进版本,主要用来解决 时间序列预测 或 需要记住过去信息的任务 ,例如:👇🏻
    • PM2.5 时间序列预测
    • 文本生成
    • 股票预测
    • 温度预测
    • 电力负载预测
  • 普通 RNN 的问题是运行久了就遗忘前面的信息 (梯度消失),而 LSTM 通过 "门结构(gates) " 让网络能够选择:👇🏻
    • 记住(Keep)
    • 忘掉(Forget)
    • 更新(Update)
  • 这些信息。
python 复制代码
# LSTM 模型定义(PyTorch)
# ---------------------------------------------------------
# 本文件实现一个简单的单层 LSTM 回归模型,用于预测下一小时 PM2.5。
# 输入维度: 24(窗口长度)
# 输出维度: 1(预测未来1小时 PM2.5)
# ---------------------------------------------------------

import torch  # 导入 PyTorch 主包,用于张量运算与设备管理
import torch.nn as nn  # 导入神经网络模块的子包,习惯性重命名为 nn

# 定义一个继承自 nn.Module 的 LSTM 模型类
class LSTMModel(nn.Module):
    def __init__(self, input_size=1, hidden_size=64, num_layers=1, dropout=0.0):
        super(LSTMModel, self).__init__()  # 调用父类构造函数,初始化模块内部状态

        # 定义一个 LSTM 层
        # input_size: 每个时间步的特征维度(这里每小时只有一个 PM2.5 值,所以是 1)
        # hidden_size: LSTM 隐藏态的维度(即每个时间步输出向量的长度)
        # num_layers: LSTM 堆叠层数(几层 LSTM 单元叠在一起)
        # batch_first=True: 输入/输出张量的形状为 (batch, seq_len, feature)
        # dropout: 当 num_layers>1 时,层间 dropout 的概率
        self.lstm = nn.LSTM(
            input_size=input_size,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True,
            dropout=dropout
        )

        # 定义一个线性全连接层,把 LSTM 的最后隐藏态映射为预测值
        # 输入维度 hidden_size -> 输出维度 1(回归预测一个数)
        self.fc = nn.Linear(hidden_size, 1)

    def forward(self, x):
        # forward 定义前向传播逻辑,x 是模型输入张量
        # 期望 x.shape = (batch_size, seq_len=24, feature=1)
        out, _ = self.lstm(x)  # 把输入传入 LSTM,out 为每个时间步的输出(shape=(batch, seq_len, hidden_size))
                               # 第二个返回值是 (h_n, c_n) ------ 最后一个时间步的隐状态与细胞状态,这里用 _ 忽略它

        # 取 LSTM 输出序列中最后一个时间步的输出作为序列级特征
        # out[:, -1, :] 的形状为 (batch_size, hidden_size)
        out = out[:, -1, :]

        # 把最后时间步的隐藏向量通过全连接层映射为标量预测值
        # 最终 out 的形状为 (batch_size, 1)
        out = self.fc(out)
        return out  # 返回预测结果(未做激活,回归任务通常直接输出实数)
相关推荐
huisheng_qaq33 分钟前
【Python基础篇-04】深入理解python的函数、参数与作用域
python·ai·函数·参数作用域
宸津-代码粉碎机5 小时前
OpenAI 连夜迎战 Grok Bot 和 Muse:AI 智能体从 “会聊天” 到 “能办事”,现在入场还来得及吗
java·大数据·人工智能·分布式·python
Y3815326625 小时前
SERP API 计费口径:credits_charged、执行失败扣费与退款边界
python·搜索引擎
蒸鱼Yuzheng7 小时前
设备端性能工件可靠导出:断点续传、哈希、manifest 与失败恢复
android·自动化测试·python·adb·数据完整性
差不多的周周7 小时前
《骨架上下文:基于零镜头骨架的动作识别的骨架侧上下文提示学习》论文核心部分详解
学习
圆圆讲门店8 小时前
挑选同城获客服务机构时需要考量的核心因素都有哪些?
大数据·网络·人工智能·python
傻啦嘿哟8 小时前
Python的默认参数把我坑惨了,原来写[]和写None的区别这么大
开发语言·python·机器学习
微小冷8 小时前
Python凸优化cvxpy初步
python·机器人·凸优化·数学规划·cvxpy
Century_Dragon8 小时前
从挂图到VR:汽修专业的一堂新能源原理课可以怎么上~
学习
在世修行8 小时前
干货:显式映射 vs 自动判据
python·插件