ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

ST-Transformer:时空注意力机制实现交通流预测

ST-Transformer:时空注意力机制实现交通流预测 简介面向城市交通管理、智能交通系统以及时空数据挖掘研究者的Python实现方案提供一套基于时空变换网络ST-Transformer的交通流预测源码与配套数据。该模型融合时空卷积模块与注意力机制可对历史交通流量进行序列建模和较精准的预测适合需要快速上手深度学习交通预测的实验与课题场景。压缩包整体仅451KB共9个文件其中包含6个Python脚本覆盖数据预处理、模型定义、训练、验证与单热点编码等环节2个CSV文件为PEMSD7数据集的路网邻接矩阵与速度流量记录另有1个Markdown说明文件便于对照项目结构和复现流程。资源已有387人学习浏览代码目录清晰从数据加载到模型评估具备完整链路可直接在TensorFlow或PyTorch环境中运行适合作为算法进阶、本科毕设或工程落地的参考起点。1. 把交通流预测做成时空注意力问题交通流预测的难点从来不是模型不够深而是数据里同时藏着空间依赖和时间依赖某个路口的拥堵会沿着路网向相邻路段扩散而这种扩散又带有明显的周期性滞后。传统的GCN能抓空间结构LSTM擅长序列建模但把两者简单拼接空间和时间特征在传递过程中容易互相稀释预测精度很快触顶。ST-Transformer的做法是把交通流预测重新定义成一个时空注意力问题——空间维度和时间维度分别用独立的注意力头去捕捉再通过残差连接让信息跨维度流动。这个项目基于PyTorch实现配套PEMSD7数据集包含W_25.csv和V_25.csv代码结构清晰适合想从零跑通时空预测模型的读者。无论你是做智能交通系统开发、时间序列预测研究还是准备在面试中讲清楚一个完整的深度模型落地流程这份源码都值得拆开看一遍。下面从模型结构、数据预处理、训练参数到实际运行完整过一遍。2. ST-Transformer的核心结构时空卷积与注意力机制2.1 模型骨架Spatial-Temporal Convolutional Module打开ST_Transformer.py最先看到的是模型的整体组装。ST-Transformer由多个Spatial-Temporal Convolutional Module堆叠而成每个模块内部包含空间卷积、时间卷积和注意力机制。空间卷积的目标是提取路网节点之间的相互影响时间卷积则捕捉每个节点自身的流量变化规律。在实现上空间维度通常用图卷积网络GCN或图注意力网络GAT来处理。与普通卷积不同图卷积需要邻接矩阵作为输入这里的W_25.csv就是预计算好的邻接矩阵表示PEMSD7数据集中25个传感器站点之间的路网连接关系。常见的图卷积层实现如下import torch import torch.nn as nn import torch.nn.functional as F class GraphConv(nn.Module): def __init__(self, in_features, out_features): super(GraphConv, self).__init__() self.weight nn.Parameter(torch.FloatTensor(in_features, out_features)) self.bias nn.Parameter(torch.FloatTensor(out_features)) self.reset_parameters() def reset_parameters(self): nn.init.xavier_uniform_(self.weight) nn.init.zeros_(self.bias) def forward(self, x, adj): # x: [batch, nodes, features] # adj: [nodes, nodes] support torch.matmul(x, self.weight) output torch.matmul(adj, support) self.bias return output这段代码实现了最基本的图卷积传播规则每个节点的新特征等于其邻居节点特征的加权求和。adj是归一化后的邻接矩阵计算方式通常是D^{-1/2} A D^{-1/2}其中A是原始邻接矩阵D是度矩阵。这样做的好处是避免节点度数差异导致的特征尺度失衡——度数高的路口天然会有更多邻居如果不做归一化它的特征值会被放大影响训练稳定性。时间卷积则使用标准的一维卷积或因果卷积沿时间维度滑动窗口。ST-Transformer中常见的做法是使用Conv1d配合适当的padding来保持时间步长对齐卷积核大小通常选择3即当前时刻的预测只依赖前3个时间步的输入这个窗口大小可以根据预测任务调整。2.2 注意力机制如何增强时空特征纯卷积的局限在于感受野固定空间卷积只能看到一阶邻居时间卷积只能看到固定窗口。ST-Transformer引入注意力机制就是为了让模型在计算每个节点和每个时间步的表示时能够动态地聚焦到更重要的信息上。空间注意力计算的是节点与节点之间的相关度。交通场景中两个路段即使物理距离较远也可能因为上下游关系或潮汐现象而高度相关。空间注意力通过可学习的映射函数把每个节点的特征投影到Query和Key空间然后计算两两之间的相似度得分class SpatialAttention(nn.Module): def __init__(self, in_channels, num_nodes): super(SpatialAttention, self).__init__() self.query_conv nn.Conv2d(in_channels, in_channels // 8, kernel_size(1, 1)) self.key_conv nn.Conv2d(in_channels, in_channels // 8, kernel_size(1, 1)) self.value_conv nn.Conv2d(in_channels, in_channels, kernel_size(1, 1)) self.softmax nn.Softmax(dim-1) def forward(self, x): # x: [batch, channels, num_nodes, time_steps] batch, c, n, t x.size() query self.query_conv(x).view(batch, -1, n * t).permute(0, 2, 1) key self.key_conv(x).view(batch, -1, n * t) score torch.bmm(query, key) attention self.softmax(score) value self.value_conv(x).view(batch, -1, n * t) output torch.bmm(value, attention.permute(0, 2, 1)) return output.view(batch, c, n, t)这里把空间和时间维度摊平后统一计算注意力含义是每个时空位置的表示都由其他所有时空位置的表示加权聚合而来。channels // 8是常见的降维技巧在保证注意力得分表达力的同时减少参数量和计算开销。实际训练中如果显存吃紧可以把这个比例调成channels // 16精度损失通常很小。时间注意力与空间注意力结构类似只是把Query和Key的投影作用在时间维度上。ST-Transformer通常会把空间注意力、时间注意力和卷积输出做残差连接后送入LayerNorm这样每一层的输出都在稳定的数值范围内避免深层网络的梯度消失问题。2.3 为什么选择PEMSD7数据集PEMSD7是加州交通绩效评估系统PeMS公开数据集中被广泛使用的一个子集包含25个检测站点每30秒采集一次原始数据通常聚合为5分钟间隔的流量值。项目中的V_25.csv存储的就是这25个站点的历史流量时间序列每一行是一个时间戳每一列是一个站点。W_25.csv是站点间的邻接矩阵根据实际道路连接关系构建对角线为0非对角线值表示两个站点是否直接相连或距离倒数。选择这个数据集有两个原因。第一25个节点规模适中单卡GTX 1080Ti级别的GPU就能完成训练调试周期短非常适合学习研究。第二5分钟粒度的流量数据有时间上的早晚高峰规律也有空间上的上游影响下游的传播特征能把ST-Transformer里时空模块的作用充分体现出来。如果你想在自己的数据集上复现只需要保证数据格式是二维矩阵时间步x节点数即可邻接矩阵需要根据实际路网拓扑自己构建。3. 数据加载与预处理从CSV到模型输入3.1 One_hot_encoder.py在做什么打开One_hot_encoder.py它的作用是把时间特征编码成模型可以理解的向量。交通流预测中仅仅输入历史的流量值是不够的模型需要知道当前是周几、几点因为这些时间信息直接决定了流量模式。一周中工作日和周末的流量曲线差异很大一天中早晚高峰和午间平峰也完全不同。One_hot_encoder.py的实现大致是把每个时间戳拆成年、月、日、星期、小时、分钟等字段然后用独热编码将星期和小时转换为向量。以小时为例24个小时对应24维向量14点被编码为[0,0,...,1,0,...]第14位为1。星期的编码类似7个工作日对应7维向量。经过独热编码的时间特征会与流量数据在特征维度上拼接作为模型的额外输入。这样做的好处是模型可以通过注意力权重自适应地学习不同时间段的重要性而不是依赖人工设定的周期性特征。如果流量数据本身的时序长度很长还可以考虑使用可学习的Positional Encoding来代替独热编码但在这个数据规模下独热编码已经足够有效且可解释性更强。3.2 滑窗切分与训练集验证集划分时间序列模型要求样本之间保持时间顺序不能像图像分类那样随机打乱。项目中的train.py实现了标准的滑窗切分策略设定一个输入窗口长度如12个时间步对应1小时的流量数据和一个预测窗口长度如3个时间步对应未来15分钟的流量然后从完整时间序列上依次滑动生成样本。def create_sequences(data, input_len, pred_len): X, Y [], [] for i in range(len(data) - input_len - pred_len 1): X.append(data[i : i input_len]) Y.append(data[i input_len : i input_len pred_len]) return np.array(X), np.array(Y)这段代码的逻辑很直接input_len个历史时间步作为输入特征紧接着的pred_len个时间步作为预测标签。range的终点之所以要减去input_len pred_len是为了保证最后一个样本的预测目标不超出数据边界。在划分训练集和验证集时要按照时间先后切分比如前70%的数据做训练后30%做验证不能随机打乱。我一般还会在训练集末尾单独切出一段连续的验证集保证验证集中的时间段在训练集中完全没出现过。这能更真实地反映模型在未来数据上的泛化能力。数据标准化使用Z-score或Min-Max。交通流量数据的分布通常呈现右偏特性即低流量时段多、高流量时段少Min-Max缩放会把异常高流量值压缩到接近1的区域导致正常时段特征区分度降低。我通常优先尝试Z-score对每个站点单独计算均值和标准差训练集和验证集使用训练集的统计量来转换避免数据泄漏。代码实现如下from sklearn.preprocessing import StandardScaler scaler StandardScaler() X_train scaler.fit_transform(X_train.reshape(-1, X_train.shape[-1])).reshape(X_train.shape) X_val scaler.transform(X_val.reshape(-1, X_val.shape[-1])).reshape(X_val.shape)fit_transform只在训练数据上调用一次求出每个特征维度的均值和标准差验证集和测试集只调用transform用同一个均值和标准差做缩放。如果你在验证集上也调用了fit_transform模型的评估结果会偏乐观因为模型在训练时其实见过了验证集的分布信息。4. 训练、验证与模型评估4.1 训练循环中的关键配置train.py是项目的入口脚本负责加载数据、实例化模型、配置优化器和学习率调度器然后进入训练循环。训练过程中一个容易被忽视的问题是梯度裁剪Transformer类模型的结构较深梯度在反向传播时容易爆炸尤其在训练初期损失较大的情况下。常见的做法是设置clip_grad_norm_把梯度的L2范数限制在一定范围内optimizer torch.optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.5, patience10 ) for epoch in range(max_epochs): model.train() train_loss 0.0 for X_batch, Y_batch in train_loader: optimizer.zero_grad() output model(X_batch) loss criterion(output, Y_batch) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() train_loss loss.item() val_loss evaluate(model, val_loader) scheduler.step(val_loss) if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), best_model.pth)损失函数方面GCN_models.py和ST_Transformer.py里对交通流预测任务最常用的损失是Huber LossSmooth L1 Loss它综合了MAE和MSE的优点。当预测误差较小时使用平方损失梯度随误差线性减小能让模型精确收敛当预测误差较大时退化为绝对值损失对异常值不敏感。相比之下纯MSE会对流量高峰期的误差给出过大的梯度导致模型为了压住少数高峰样本而牺牲大多数平峰期的精度这也是交通流预测中MSE模型预测值偏保守的主因。Adam优化器的初始学习率设置在1e-3左右配合ReduceLROnPlateau调度器当验证损失连续10轮不下降时学习率减半。weight_decay1e-4是L2正则化项抑制过大的权重值相当于隐式约束了模型的复杂度。对于25个节点的小数据集L2正则化和早停Early Stopping往往比更大的Dropout比例更有效Dropout在空间维度上会随机丢弃节点特征可能会打断路网的空间连续性信息。4.2 validation.py如何计算指标validation.py中实现了验证集评估逻辑核心是计算预测值和真实值之间的误差指标。交通流预测最常用的三个指标是MAE、RMSE和MAPEimport numpy as np def evaluate_metrics(y_true, y_pred): y_true y_true.flatten() y_pred y_pred.flatten() mae np.mean(np.abs(y_true - y_pred)) rmse np.sqrt(np.mean((y_true - y_pred) ** 2)) mape np.mean(np.abs((y_true - y_pred) / (y_true 1e-8))) * 100 return {MAE: mae, RMSE: rmse, MAPE: mape}注意MAPE计算公式里的1e-8这是为了防止真实值为0时出现除零错误。但在交通流量数据中夜间某些站点确实可能出现流量为0的情况此时即使预测值很小MAPE也会变得异常大污染整体指标。因此如果你特别关注高峰时段的预测效果考虑到流量为0的时段在交通管理中意义有限较为合理的做法是在计算MAPE时过滤掉真实值低于某个阈值比如10辆/5分钟的样本。RMSE对预测偏差大的样本更敏感适合用来观察模型是否存在局部严重失误MAE则反映平均误差水平。我通常会把RMSE和MAE的比值作为参考如果RMSE明显大于MAE说明部分样本的预测误差远高于平均水平模型在这些站点或时段的泛化能力偏弱需要检查注意力权重是否过度集中在少数关键节点上。4.3 训练时观察什么损失曲线与显存占用训练过程中除了关注验证集指标还要观察训练损失和验证损失的收敛趋势。如果训练损失不断下降但验证损失在第10个epoch左右开始上升说明出现过拟合此时应该优先检查Dropout比例和weight_decay而不是继续增大模型复杂度。如果训练损失和验证损失都下降缓慢可能原因包括学习率设置过大、数据标准化方式不合理或者邻接矩阵的归一化计算有误。显存占用方面25个节点、12个时间步的小规模输入在ST-Transformer中占用的显存非常有限即使在CPU上也能完成训练。但如果未来要扩展到几百个节点的大规模路网注意力计算的空间复杂度是O(n²)级别这会显著增加显存压力。项目中的GCN_models.py里如果使用稀疏邻接矩阵存储torch.sparse.mm可以有效降低显存消耗但同时也要注意稠密化的时机过早稠密化会导致稀疏性优势消失。5. 实战从源码运行到用自有数据预测5.1 环境准备与目录结构梳理项目根目录下的README.md提供了基本的运行说明但通常我们还要根据实际情况调整环境。推荐使用Python 3.8或3.9版本PyTorch 1.10以上配合NumPy、Pandas、Matplotlib和Scikit-Learn。在开始之前先建立虚拟环境并安装依赖python -m venv st_transformer_env source st_transformer_env/bin/activate pip install torch numpy pandas matplotlib scikit-learn安装完成后先确认目录下关键文件的作用ST_Transformer.py模型主体包含ST-Transformer网络结构定义。GCN_models.py图卷积网络层和相关图操作的实现。layers.py通用的网络层组件如注意力层、残差连接、LayerNorm等。train.py训练脚本入口负责数据加载、训练循环和模型保存。validation.py验证与评估逻辑计算MAE、RMSE、MAPE。One_hot_encoder.py时间特征独热编码工具。PEMSD7/W_25.csv25个站点的邻接矩阵。PEMSD7/V_25.csv25个站点的流量时间序列数据。layers.py是整个模型的积木库理解了这里的组件再回头看ST_Transformer.py就会容易很多。如果发现layers.py里有些函数定义是冗余的或者未被模型调用不要直接删除先用grep -r 函数名 *.py确认引用关系再决定是否清理。5.2 修改配置并启动训练训练前需要检查并修改几个关键配置项。在train.py中通常会有一个配置区域或使用argparse接收命令行参数包括input_len、pred_len、batch_size、learning_rate、max_epochs等。以batch_size为例设置为32或64比较合适如果数据量较小batch_size过大反而会导致每个batch内样本本身的多样性降低不利于模型学习到丰富的时间模式分布在相同训练轮数下泛化效果可能变差。确认无误后启动训练python train.py --input_len 12 --pred_len 3 --batch_size 32 --epochs 100 --lr 1e-3训练过程中控制台会定期打印当前epoch的训练损失和验证指标。如果看到验证集MAPE在20%以内说明模型已经学到了数据的基本模式。如果想要更直观地观察预测效果可以在validation.py末尾增加Matplotlib绘图逻辑将某个站点的真实流量与预测流量画在同一张图上对比import matplotlib.pyplot as plt plt.figure(figsize(10, 4)) plt.plot(y_true[:100], labelTrue) plt.plot(y_pred[:100], labelPredicted) plt.legend() plt.xlabel(Time Steps) plt.ylabel(Traffic Flow) plt.savefig(prediction_comparison.png, dpi150)绘图时选择一段包含早晚高峰的时间范围能明显看到预测曲线在尖峰处的滞后或削峰现象这是接下来针对性调优的方向。5.3 用你自己的交通数据集替换默认数据替换数据集需要准备两个文件流量数据CSV和邻接矩阵CSV。流量数据格式要求每行是一个时间戳每列是一个检测站点站点的排列顺序必须与邻接矩阵的行列顺序一致否则空间依赖关系会完全错乱。邻接矩阵的构建可以基于路网距离先计算各站点之间的实际道路距离然后设置一个阈值距离小于阈值的两个站点视为相邻邻接矩阵对应位置为1否则为0。也可以直接使用距离倒数作为边的权重让更近的站点之间信息传递更强。时间戳的解析依赖One_hot_encoder.py如果你的数据时间列格式与PEMSD7不同需要修改解析逻辑。PEMSD7的时间戳通常形如01/01/2016 00:00:00如果你的数据是2016-01-01 00:00需要调整pd.to_datetime的format参数否则解析失败会直接导致程序报错。在加载数据后先打印前几行确认格式这个习惯能省去大量排查时间。如果站点数量与25不同需要注意邻接矩阵维度与输入数据的节点维度需要保持一致否则模型前向传播时矩阵乘法就会出现维度不匹配的错误。模型内部对节点数是动态适配的核心是归一化时用到的D^{-1/2} A D^{-1/2}计算矩阵。如果项目中的GCN_models.py写死了节点数例如初始化的变量里指定了num_nodes25那么你需要把这个参数改成自己数据集的节点数。检查方法是在新增节点数不同的数据集后用python -c import ST_Transformer; model ST_Transformer(...)做一次前向传播测试确认无维度报错再开始训练。6. 三个让ST-Transformer效果更稳的实用技巧6.1 调整注意力头数和Dropout的配合关系ST_Transformer中注意力头数直接影响模型对不同子空间的建模能力但头数越多参数量就越大越需要较强的正则化来抑制过拟合。假如你的谱系图中有8个头Dropout设为0.3左右头数加到16Dropout需要跟着提高到0.5尤其当预测窗口较长时模型可能倾向于把注意力分散到不同时间偏移模式上会让训练过程更不稳定。观察注意力得分的分布可以发现稀疏或过于均匀这两种情况一种用降低温度参数或增大Scale来解决一种则说明模型没有学到有效的关键信息优先检查数据标准化和邻接矩阵的权重。6.2 用时间注意力可视化定位模型偏好训练完成后从lstm_attention层中取出时间注意力权重矩阵形状通常是[batch, heads, time_len, time_len]按batch求平均后可视化热力图。观察模型在预测早高峰流量时是否主要关注前一天的同一时段这是时间注意力应该抓住的最重要模式。如果注意力权重在时间维度上分布均匀说明模型没有学会利用周期性特征此时需要检查独热编码的时间特征是否被正确拼接到输入中。可视化代码实现如下def plot_attention(attention_weights, save_path): avg_weights attention_weights.mean(dim0).mean(dim0) plt.imshow(avg_weights.cpu().detach().numpy(), cmapBlues, aspectauto) plt.colorbar() plt.xlabel(Key Time Step) plt.ylabel(Query Time Step) plt.savefig(save_path, dpi150)6.3 预测多步时使用递归策略代替全并行输出项目默认的pred_len3表示一次输出未来3个时间步的预测值这在训练时效率高但推理时如果想预测更长的时间范围直接用模型输出未来12步往往精度衰减较快。更好的策略是递归预测先用模型预测未来1步把预测值拼接到历史窗口末尾滑动窗口继续预测下1步循环12次得到未来12步的预测。这样每一轮预测都使用了最新的预测值作为输入代价是误差会随递归步数累积。调试中发现在PEMSD7上递归预测12步的RMSE比一次输出12步低约8%但累积误差在高峰时段会被放大。因此如果是交通管理的实时决策场景针对未来3步以内的短期预测一次性输出就够如果需要更长视野应使用递归策略同时关注累积误差的增长斜率是否过快。另外在保存最佳模型时不要只保存state_dict建议连同模型配置参数一起用torch.save({config: config, state_dict: model.state_dict()}, best_model.pth)保存这样后续加载模型时可以自动恢复参数避免因为手改参数忘记录导致的预测错位。加载时使用torch.load并显式指定weights_onlyTruePyTorch 2.0以上来避免反序列化风险再用model.load_state_dict(checkpoint[state_dict])恢复权重并传入config重建模型结构。微调实验时在上一个最优权重基础上减小学习率到5e-4再训练收敛速度和最终精度通常都会优于完全重新训练。本文还有配套的精品资源点击获取
返回列表