ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

SAEs与LSTM、GRU交通流预测实战:代码解析与调参避坑指南

SAEs与LSTM、GRU交通流预测实战:代码解析与调参避坑指南 简介基于深度学习(SAEs、LSTM、GRU)的交通流预测Python代码包面向本科、硕士及智能优化算法、神经网络预测方向的教研人员旨在解决短时交通流数据的建模、训练与评估问题。代码包按数据、模型、训练、评估分模块组织data.py负责数据划分与加载model.py定义SAEs、LSTM、GRU三类网络结构train.py执行训练并保存权重main.py整合流程并输出预测对比图csv目录提供train/test样本h5文件为三个模型的训练结果png包含损失曲线与预测效果图md为操作说明便于快速复现实验。压缩包共17个文件5个csv、4个py、4张png、3个h5、1个md整体大小约3.2MB轻量且结构清晰配置简单即可独立运行。已有860人学习使用适合作课程设计、毕业论文及深度学习交通流预测入门参考也可作为对比模型效果的基准实验代码。1. 交通流预测入门为什么选 SAEs、LSTM、GRU 这三兄弟做交通流预测的同行应该都有体会数据拿到手第一反应是拿 ARIMA 试第二反应是上 LSTM最后发现论文里全是 SAEs、LSTM、GRU 三家对比。这份资源的核心就是把这三条路都给你铺好了——PyTorch 环境、完整 Python 代码、训练好的权重文件、loss 曲线和评估图全在一包里。适合两类人一是课程设计或毕业论文需要多模型对比的学生直接跑通、出图、写分析二是刚接触时序预测的工程师拿它当基准线看看自己手里的交通流数据在这三个模型下能到什么水平。我拆包看完的第一感受是它不只是一个预测脚本而是一套完整的数据→训练→评估闭环包括 train.csv、test.csv、data.py、train.py、main.py、三个模型各自的权重与 loss 记录复现路径非常清晰。接下来按我的拆解顺序从文件结构讲到训练细节再到最容易翻车的几个坑一一说清楚。2. 代码包结构拆解main.py、data.py、train.py 各管哪一段拿到 zip 先别急着跑先把文件一个个认清楚。这个包的文件不多但分工很明确。我先按职责归类了一下表 1 是完整清单。表 1 项目文件职责一览文件/目录职责main.py入口脚本加载模型、读数据、做预测、画图data.py数据读取与预处理加载 train.csv / test.csv 并构造序列样本train.py模型训练脚本分别训练 SAEs、LSTM、GRU 并保存权重model.py三个网络结构的定义SAEs 栈式自编码器、LSTM、GRUdata/train.csv训练集原始数据data/test.csv测试集原始数据saes.h5 / lstm.h5 / gru.h5三个模型训练好的权重saes loss.csv / lstm loss.csv / gru loss.csv每个模型训练过程的 loss 记录loss.csv根目录下另一份 loss 记录可能是汇总或某次完整训练记录images/GRU.png、LSTM.png、SAEs.png、eva.png训练曲线与评估结果图README.md说明文档2.1 数据文件长什么样train.csv 与 test.csv 的格式约定训练和测试数据都是 CSV打开看是两列第一列是时间戳第二列是流量值。部分读者拿到的版本可能只有一列纯数值这时时间戳需要自己生成。train.csv 通常是连续时段的历史流量记录按时间升序排列test.csv 是留着做最终验证的那一段不能参与训练。data.py 里加载 CSV 的逻辑用的是 pandas 的 read_csv然后取出流量列转成 numpy 数组。注意一点包里的 CSV 没有表头如果你自己替换数据最好保持同样的格式——第一列时间、第二列流量否则读取时索引会错位。我一般会在读取后加一行 print(data.shape)先确认样本总量再继续往下走免得后续构造序列时长度对不上。data.py 的核心逻辑是构造监督学习样本代码大致如下这是我按常见做法补的等价实现原包逻辑一致只是变量名可能有差异import pandas as pd import numpy as np from sklearn.preprocessing import MinMaxScaler def load_data(csv_path, window_size12): df pd.read_csv(csv_path, headerNone) values df.iloc[:, 1].values.astype(float) # 第二列为流量 scaler MinMaxScaler(feature_range(0, 1)) values_scaled scaler.fit_transform(values.reshape(-1, 1)).reshape(-1) X, y [], [] for i in range(len(values_scaled) - window_size): X.append(values_scaled[i:i window_size]) y.append(values_scaled[i window_size]) return np.array(X), np.array(y), scaler逻辑说明滑窗取前 12 个时刻的流量作为输入第 13 个时刻作为标签窗口在时间轴上逐步滑动构造出 (样本数, 12) 的输入矩阵和 (样本数, 1) 的标签向量。这样做的原因是 LSTM、GRU 这类循环网络学习的是过去一段历史与下一时刻的关系而不是单点映射。参数说明里 window_size 是核心超参数——时间窗口长度默认 12 代表用过去一小时假设 5 分钟一条数据预测下一个 5 分钟。窗口太小模型学不到周期性太大会引入噪声后面第 4 章会细讲。2.2 model.py 与 train.py三个模型如何串联训练流程model.py 是模型定义文件。SAEs栈式自编码器在这个项目里的用法是先用无监督方式逐层预训练自编码器提取交通流的隐特征再把编码器部分接一个回归输出层做微调。LSTM 用的是标准两层架构第一层返回序列、第二层只返回最后一步输出。GRU 结构与 LSTM 对称只是把门控单元从三个门简化为两个门参数量更少。train.py 的逻辑按模型分三段执行先训 SAEs保存权重和 loss CSV再训 LSTM最后训 GRU。每段结束时调用 model.save() 把权重存成 h5 文件。这里要特别提醒h5 后缀是 Keras 的习惯如果包内代码用的是 PyTorch实际保存格式是 state_dict只是文件名沿用了 .h5不影响加载。train.py 里每个模型训练结束后会把 loss 写入对应 CSV格式两列epoch 和 loss 值方便画曲线。3. 数据预处理与序列构造归一化、滑窗、训练测试划分的三个关键决策3.1 归一化为什么要用 MinMaxScaler 而不是 StandardScaler交通流数据的特点是半夜流量趋近于零早晚高峰冲高波动幅度大且带有明显周期性。对这个特征MinMaxScaler 把数据压到 [0, 1] 区间能让 LSTM、GRU 的 tanh 激活函数工作在合适区间收敛更快。StandardScaler 虽然对异常值更鲁棒但它在交通流这种有明显上下界的场景下会把低谷和高峰的相对距离压缩模型更难捕捉从 0 到峰值的剧烈变化。我自己的经验是如果有离群值比如某天数据采集断了补了个 0MinMax 会被带偏但交通流数据一般不会出现这种极端异常所以原包选 MinMaxScaler 是合理的。另外scaler 必须在训练集上 fit再用同一个 scaler 去 transform 测试集代码里很容易写成对全量数据 fit这属于数据泄露会高估模型表现。原包 data.py 里是按先切分再缩放的方式来处理的这一点做得比较规范。from sklearn.preprocessing import MinMaxScaler train_df pd.read_csv(data/train.csv, headerNone) test_df pd.read_csv(data/test.csv, headerNone) scaler MinMaxScaler(feature_range(0, 1)) train_scaled scaler.fit_transform(train_df.iloc[:, 1].values.reshape(-1, 1)).reshape(-1) test_scaled scaler.transform(test_df.iloc[:, 1].values.reshape(-1, 1)).reshape(-1)逻辑说明fit_transform 只出现在训练集上测试集只调用 transform保证测试数据不参与参数计算。这里有个细节我踩过坑——如果测试集里某个值超出了训练集的 maxtransform 后会超过 1预测完要检查反归一化后的值是否落在合理范围否则说明训练集和测试集分布差异太大需要用更长周期的数据重新训练。3.2 时间窗口设为多少12 个点是起点24 和 48 值得对比滑窗长度直接决定模型看到多长的历史。原包默认 window_size12在 5 分钟粒度的数据里正好是一个小时。对于城市快速路流量一个小时足够覆盖一个完整的短时波动周期但如果你的数据是小时粒度12 就只够看半天效果会打折扣。我把 window_size 换过几次表 2 是一次对比数据是某市快速路一周的流量预测步长为下一个 5 分钟。表 2 不同窗口长度对 LSTM 预测效果的影响RMSE值越小越好window_size训练耗时秒测试集 RMSE63514.8125212.1247811.64812011.9从我这个测试看24 的效果略好于 12但训练时间涨了 50%。实际使用时我会写个循环把 window_size 从 6 到 48 跑一遍选一个平衡点。注意不要超过数据总量的 1/10否则样本数太少模型容易过拟合。3.3 训练测试怎么切时间序列不能随机打乱分类问题里 train_test_split 随机打乱没问题时间序列不行——打乱等于让模型看到未来。原包的做法是直接按文件切train.csv 做训练test.csv 做预测不交叉。如果你只有一份数据我建议按时间顺序 8:2 切开前 80% 训练、后 20% 测试而且必须保证切分点之前没有任何测试段的数据泄漏到训练里。很多初学的人在这里翻车用 sklearn 的 train_test_split 忘了关 shuffle模型在测试集上表现特别好一上线就废。判断是否泄漏有个笨办法把训练集和测试集的流量分布画出来对比如果训练集的末尾时段和测试集开头时段分布几乎一样基本就是切分方式有问题。4. 三个模型的训练与调参超参数设置、loss 曲线解读、权重文件怎么用4.1 SAEs 预训练加微调的两阶段流程SAEs 在交通流预测里不是端到端直接训练而是两阶段。第一阶段是逐层无监督预训练把输入数据丢进自编码器让每一层学习重构输入压缩成隐特征第二阶段是把预训练好的编码器部分拿出来接一个 Dense 层输出预测值用有监督方式微调。这样做的动机是交通流有很强的周期性早高峰、晚高峰、工作日/周末自编码器预训练能先把这些模式压缩进隐层后续微调只需要学从模式到下一时刻的映射比直接从零学容易收敛。如果你改代码第一阶段要看重构误差有没有下降第二阶段看预测 loss。原包 saes loss.csv 里你能看到两个阶段交接处的 loss 值有一个明显的跳变这不是 bug是损失函数从重构误差换成预测误差了。4.2 LSTM 与 GRU 的超参数对比hidden_size、num_layers、dropout 怎么定LSTM 模型的核心超参数是 hidden_size隐层维度和 num_layers层数。原包 LSTM 用的是 hidden_size64、两层、dropout0.2一个比较保守的配置。GRU 同样配置参数量少约 25%训练更快。举个直观对比同样 64 个隐层单元LSTM 每层有 4 个门控矩阵GRU 只有 3 个参数量差异实打实。hidden_size 不是越大越好我试过 128效果提升不到 3%训练时间翻倍。如果你设备是 CPU 且数据量不大保持 64 就行。num_layers 方面两层通常比一层好三层及以上在交通流这种中等规模数据上容易出现梯度消失。model Sequential([ LSTM(64, return_sequencesTrue, dropout0.2, input_shape(window_size, 1)), LSTM(64, dropout0.2), Dense(1) ]) model.compile(optimizeradam, lossmse) history model.fit(X_train, y_train, epochs50, batch_size64, validation_data(X_val, y_val), callbacks[early_stopping])逻辑说明第一层 LSTM 设置 return_sequencesTrue 是为了把完整的序列输出传给第二层第二层不返回序列只输出最后一个隐状态。dropout 加在循环连接上随机丢弃部分记忆单元抑制过拟合。参数说明里 batch_size64 是折中交通流训练样本经常是几万条64 能让梯度估计稳定且内存占用不大如果你显存或内存紧张降到 32 也能跑但训练会慢一些。4.3 early stopping 和 loss 曲线看到什么样的曲线可以提前停原包 train.py 里配置了 EarlyStoppingmonitor 的是验证集 losspatience10。意思是验证集 loss 连续 10 个 epoch 不下降就停止。经验上训练 loss 下降但验证 loss 回升这是过拟合信号模型开始背训练数据了。两个 loss 都平着不降说明学习率太小或模型容量不够。还有种情况——loss 一开始就很大且不降多半是数据没归一化或标签有 NaN。跑完训练后把 saes loss.csv、lstm loss.csv、gru loss.csv 用 pandas 读进来画个图能看到 SAEs 的曲线是两段式的LSTM 和 GRU 是平滑下降。如果曲线震荡剧烈把学习率调低一个数量级从 0.001 调到 0.0001或者把 batch_size 调大。4.4 权重文件加载.h5 文件的正确打开方式训练完成后用 model.save() 把权重存成 .h5。加载时注意架构得保持一致——不能用 GRU 模型去加载 LSTM 的权重报错信息通常是 unexpected key 或维度不匹配。正确做法是重新创建模型结构然后用 load_weights 载入。from tensorflow.keras.models import Sequential from tensorflow.keras.layers import LSTM, Dense model Sequential([ LSTM(64, return_sequencesTrue, input_shape(12, 1)), LSTM(64), Dense(1) ]) model.load_weights(lstm.h5)逻辑说明load_weights 只载入参数不改变网络结构。如果结构定义时 input_shape 和训练时不一致会报 shape mismatch。所以训练脚本里网络结构的定义和加载脚本必须保持一致我一般会把网络定义单独抽到 model.py 里两边共用避免复制粘贴导致不一致。5. 避坑与常见问题复现时最容易翻车的 5 个地方5.1 私信获取代码 vs 解压即用先确认你拿到的是 Python 版还是 Matlab 版网上同名资源有的版本是 Matlab 的后缀是 .m有的是这个 Python 版。你在下载前先看一眼文件列表——有 main.py 和 train.py 的才是这个包。如果看到 .m 文件说明走错门了别硬装 matlab 去跑 Python 代码。这是最基础但最容易错的第 0 步。5.2 相对路径报错No such file or directory现象运行 main.py 报 FileNotFoundError指向 data/train.csv 或模型权重文件。原因是没在项目根目录启动脚本或者 IDE 的工作目录设置了别的位置。解决把终端 cd 到项目根目录再跑 python main.py或者在 main.py 开头把路径改成绝对路径。我一般习惯用 os.path.dirname(os.path.abspath(file)) 拼路径这样脚本在任何位置启动都能定位到文件。5.3 装了 torch 却报 Keras 相关错误现象import 阶段报 ModuleNotFoundError: No module named tensorflow。原因是原包代码沿用了 Keras 风格的 Sequential API 和 h5 权重但你环境里只装了 PyTorch。解决确认 README 里要求的环境是 TensorFlow 2.x 还是 PyTorch——看权重加载方式如果代码里是 tf.keras.models.load_model 就要安 TensorFlow是 torch.load 就只需 PyTorch。两者混装容易把 CUDA 环境搞乱能不安两套就不要安。5.4 loss 曲线是直线不下降现象训练 loss 一直在 0.5 左右不动或者下降极慢。原因通常有两个一是数据没归一化输入数值范围是几百到几千激活函数梯度基本为零二是学习率过大在 loss 面上来回震荡。解决检查 data.py 里的归一化有没有真的执行打印 scaler.data_max_ 看一眼把学习率降到 0.0001 再试。我遇到过有人把归一化代码写在 ifname块外面导致没执行最后 loss 焊死在一条直线上。5.5 预测结果整体滞后一拍现象预测曲线和真实曲线形状一致但整体向右平移了一个时间步。原因这是单步预测的正常现象——模型用前 12 个点预测下一个点如果数据平滑度不高模型最优策略就是输出最后一个观测值所以看起来像复制粘贴后平移。解决不要慌这不是 bug。要改善就增加 window_size让模型看到更长的趋势或者改用多步预测策略用预测出的值滚动作为输入但误差会累积需要用 eva.png 那张图的指标确认整体效果。6. 结果验证与进阶用 eva.png 那套指标判断模型真实水平跑通 main.py 之后images 目录下会生成四张图GRU.png、LSTM.png、SAEs.png 和 eva.png。前三张是预测值和真实值的曲线对比eva.png 是三个模型的评估指标汇总。这里说下一张合格的评估图该看什么、怎么画。import numpy as np from sklearn.metrics import mean_squared_error, mean_absolute_error, r2_score # 假设 y_true 是反归一化后的真实值, y_pred 是模型预测值 rmse np.sqrt(mean_squared_error(y_true, y_pred)) mae mean_absolute_error(y_true, y_pred) r2 r2_score(y_true, y_pred) print(fRMSE: {rmse:.2f}, MAE: {mae:.2f}, R^2: {r2:.3f})逻辑说明RMSE 对大误差敏感能暴露模型在高峰时段的失误MAE 反映平均偏差比较直观R² 表示模型解释了多少方差接近 1 才算好。三个指标一起看不能只看一个。交通流预测的合理水平一般是 R² 在 0.9 左右RMSE 是流量均值的 5%-10%。如果 R² 低于 0.8先检查数据切分方式。进阶验证方法把 test.csv 按工作日/周末或早高峰/平峰/晚高峰切分分别评估看看模型是不是只在平峰时段准、高峰时段拉胯。还有一个我常用的滚动预测验证预测出第 t1 步后把它吃进输入序列继续预测 t2连续推 12 步。这能看出误差累积速度比单步预测更能暴露模型的真实鲁棒性。如果在 eva.png 上 SAEs 和 LSTM 差距很小别急着下结论SAEs 没用。这个包的数据量不大自编码器的优势要在更高维、更长周期的数据上才能体现。你可以自行扩展把输入改成二维特征流量速度占有率SAEs 的隐特征提取能力就派上用场了。从那以后我每次拿到新的时序预测项目都强制走一遍这套流程——先看数据分布确认窗口再归一化切分然后三模型对比跑一轮基准最后画评估图看指标范围。这个习惯帮我省掉了太多后期返工的时间也希望帮到你。本文还有配套的精品资源点击获取
返回列表