ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

交通流预测中的ST-Transformer原理与实战

交通流预测中的ST-Transformer原理与实战 简介本资源是一套基于Python实现的交通流预测高分项目源码与配套数据集面向计算机、交通工程或人工智能方向的本科生及研究生适用于毕业设计、课程设计与期末大作业等实践场景。项目采用时空变换网络ST-Transformer建模方法融合图卷积与时间注意力机制可有效捕捉路网拓扑结构与动态交通时序特征。压缩包共8个文件含6个核心Python脚本如ST_Transformer.py、train.py、GCN_models.py等覆盖模型构建、训练验证与数据编码及2个CSV格式真实交通流数据W_25.csv与V_25.csv整体仅484KB轻量易部署。已有171人学习下载代码全程中文注释详尽逻辑清晰、模块解耦合理附带PEMSD7标准数据预处理流程与One-hot编码工具新手可快速理解模型架构与训练全流程具备完整复现能力与工程迁移价值。1. 用 Python 跑通交通流预测的时空变换网络不是调个库就完事——它要同时建模“哪里堵”和“什么时候堵”你下载了一个叫Python实现用于交通流预测的时空变换网络源码数据集高分项目.zip的压缩包解压后看到model.py、data_loader.py、train.py和一个data/文件夹。但直接python train.py却报错ModuleNotFoundError: No module named torch装完 PyTorch 又卡在KeyError: sensor_102再查data/下只有PeMSD7_W_228.csv和PeMSD7_V_228.csv—— 这根本不是你本地路口的传感器编号。这说明交通流预测中的时空变换网络ST-Transformer不是端到端黑盒它强依赖特定结构的路网拓扑与时间序列对齐方式。它解决的是城市级主干道断面流量的 15–60 分钟超前预测问题核心难点在于——同一时刻A 路口车流激增可能由 B 路口 3 分钟前的事故引发时间滞后而 C 路口虽物理距离远却因平行快速路共享车流空间非欧。因此本项目不是教你怎么用sklearn做回归而是带你从原始传感器时序中重建路网图结构、构造时空注意力掩码、对齐多尺度周期性早高峰 vs 周末夜宵流最终让模型真正“看懂”城市脉搏的跳动节律。适合已掌握 PyTorch 基础、接触过图神经网络或时间序列建模且手头有真实地磁/线圈检测器数据的交通工程、智能网联或城市计算方向从业者。2. 为什么必须用 ST-Transformer 而非 LSTM 或 GCN——交通流的时空耦合不可拆解2.1 传统方法在交通场景下的三重失效交通流数据天然具备强周期性日周期、周周期、突发性事故、天气和空间依赖性上下游、平行路网。LSTM 类模型仅建模时间维度会将 A 路口上午 8:00 的高流量与 B 路口下午 14:00 的高流量错误关联GCN 类模型虽引入图结构但其邻接矩阵固定如仅按地理距离定义无法表达“早高峰时 A→B 的车流强度是晚高峰的 3.2 倍”这类动态空间关系而简单 CNN 则完全忽略传感器间的拓扑约束把路网当成图像像素处理导致预测结果在交叉口处出现物理不一致如上游流量为 200 辆/小时下游反推为 80 辆/小时。提示PeMSD7 数据集中的W文件是路段拓扑权重矩阵Weight matrixV是流量观测值Volume二者必须严格对齐。若你用自己的数据不能直接删掉W文件改用全连接——那等于放弃空间建模。2.2 ST-Transformer 的双通道注意力设计原理ST-Transformer 的核心创新在于将时空建模解耦为两个并行子网络再通过门控机制融合Spatial Transformer Block输入为(B, T, N, C)其中B是 batch sizeT是时间步如 12×5min1hN是传感器数量如 228C是特征维度流量、速度、占有率。它不使用预设邻接矩阵而是让模型自己学习每对传感器(i, j)在当前时间窗口内的动态相关性# spatial_attn.py 中关键计算简化 Q_s self.W_q_s(x).view(B, T, N, self.d_k) # [B,T,N,d_k] K_s self.W_k_s(x).view(B, T, N, self.d_k) V_s self.W_v_s(x).view(B, T, N, self.d_k) attn_weights torch.softmax(torch.einsum(btik,btjk-btij, Q_s, K_s) / np.sqrt(self.d_k), dim-1) # 注意这里 attn_weights.shape (B, T, N, N)每个时间步 t 都有独立的空间注意力图Temporal Transformer Block将每个传感器i的历史序列(B, T, C)视为一个 token 序列学习跨时间步的长期依赖# temporal_attn.py 中位置编码处理 pos_encoding torch.zeros(T, C) pos torch.arange(0, T, dtypetorch.float).unsqueeze(1) # [T,1] div_term torch.exp(torch.arange(0, C, 2).float() * (-np.log(10000.0) / C)) pos_encoding[:, 0::2] torch.sin(pos * div_term) # 偶数位正弦 pos_encoding[:, 1::2] torch.cos(pos * div_term) # 奇数位余弦 x x pos_encoding.unsqueeze(0) # [B,T,C] [1,T,C]2.3 为什么必须重写数据加载器——原始 CSV 不等于模型输入张量项目中的data_loader.py并非通用读取器它执行三个不可跳过的转换缺失值插补交通传感器常因断电/通信故障丢失整段数据。直接用pandas.fillna(methodffill)会导致早高峰突增被平滑成缓坡。本项目采用ST-MIDAS 插补法代码见utils/impute.py先用 GCN 学习空间相似性再用 LSTM 学习时间模式联合优化缺失值周期性特征工程除原始流量外必须注入day_of_week、hour_of_day、is_holiday三类周期信号并做 sin/cos 编码避免 0 与 23 小时的语义断裂滑动窗口切片将连续时序(total_len, N)切为(num_samples, input_len, N, C)其中input_len121h 输入output_len315min 预测。关键参数seq_len12必须与模型中TemporalTransformer的最大位置编码长度一致否则pos_encoding索引越界。# 验证数据加载是否正确在 train.py 开头插入 loader DataLoaders(data_dirdata/, batch_size32, input_len12, output_len3) x, y next(iter(loader.train_loader)) print(fInput shape: {x.shape}) # 应输出 torch.Size([32, 12, 228, 3]) print(fOutput shape: {y.shape}) # 应输出 torch.Size([32, 3, 228, 1]) print(fFeature dim: {x.shape[-1]}) # 第 3 维必须为 3[flow, speed, occupancy]3. 从零复现 ST-Transformer 模型结构——逐层解析model.py的 7 个核心组件3.1 整体架构Encoder-Decoder 结构的交通专用改造原始 Transformer 的 Encoder-Decoder 用于机器翻译而交通预测是序列到序列的映射故本项目采用单 Encoder 线性 Decoder设计非自回归Encoder堆叠L3层 Spatial Temporal Block交替排列Decoder仅用一层全连接层将 Encoder 输出(B, T, N, d_model)映射为(B, output_len, N, 1)关键区别Encoder 输入是(B, input_len, N, C)但输出需重排为(B, N, input_len, d_model)才能送入 Temporal Block 处理时间维度。# model.py 中 Encoder 定义精简 class STTransformerEncoder(nn.Module): def __init__(self, num_layers3, d_model64, n_heads4, d_ff256, dropout0.1): super().__init__() self.layers nn.ModuleList([ STTransformerBlock(d_model, n_heads, d_ff, dropout) for _ in range(num_layers) ]) self.norm nn.LayerNorm(d_model) def forward(self, x, spatial_maskNone, temporal_maskNone): # x: [B, T, N, d_model] → 经过 L 层后仍保持此形状 for layer in self.layers: x layer(x, spatial_mask, temporal_mask) return self.norm(x)3.2 Spatial Transformer Block动态图学习的实现细节该模块不依赖预定义邻接矩阵而是通过传感器嵌入Sensor Embedding初始化节点表征再用自注意力计算动态边权# spatial_block.py 中 SensorEmbedding 实现 class SensorEmbedding(nn.Module): def __init__(self, num_sensors, embed_dim): super().__init__() self.embedding nn.Embedding(num_sensors, embed_dim) # PeMSD7 有 228 个传感器故 num_sensors228 self.pos_embedding nn.Parameter(torch.randn(1, 1, num_sensors, embed_dim)) def forward(self, x): # x: [B, T, N, C] → 仅用 N 维度索引 embedding sensor_emb self.embedding.weight # [N, embed_dim] pos_emb self.pos_embedding # [1, 1, N, embed_dim] return sensor_emb pos_emb # [1, 1, N, embed_dim]注意SensorEmbedding的pos_embedding是可学习参数而非正弦位置编码——因为传感器物理位置固定但其功能角色如“主干道入口”vs“支路分流点”需由模型自主发现。3.3 Temporal Transformer Block如何处理不同长度的时间依赖交通流存在多尺度周期5min 微观波动、1h 宏观趋势、24h 日周期。本项目采用Multi-Scale Temporal Attention将输入时间序列(B, N, T, C)拆分为 3 组T12→[12, 6, 3]对应 1h、30min、15min 窗口每组独立进 Temporal Transformer输出后拼接最终用卷积层融合多尺度特征。# temporal_block.py 中 multi-scale 实现 class MultiScaleTemporalAttention(nn.Module): def __init__(self, d_model, scales[12,6,3]): super().__init__() self.scales scales self.attns nn.ModuleList([ TemporalSelfAttention(d_model, scale) for scale in scales ]) self.conv_fuse nn.Conv2d(len(scales), 1, kernel_size1) def forward(self, x): # x: [B, N, T, d_model] feats [] for i, scale in enumerate(self.scales): # 截取最后 scale 个时间步 x_scale x[:, :, -scale:, :] # [B, N, scale, d_model] feat self.attns[i](x_scale) # [B, N, scale, d_model] feats.append(feat.unsqueeze(1)) # [B, 1, N, scale, d_model] # 拼接后融合 fused torch.cat(feats, dim1) # [B, 3, N, scale_max, d_model] out self.conv_fuse(fused.permute(0,1,4,2,3)).squeeze(1) # [B, d_model, N, scale_max] return out.permute(0,2,3,1) # [B, N, scale_max, d_model]3.4 损失函数与训练策略MAE 之外必须加物理约束单纯用 MAE 或 MSE 会导致预测值违反交通流守恒律如某路口流入量持续大于流出量。本项目在损失中加入Traffic Flow Conservation Loss# loss.py 中物理约束实现 def traffic_conservation_loss(pred, true, adj_matrix, alpha0.2): # adj_matrix: [N, N], 行表示流入列表示流出需根据路网方向构建 # pred: [B, T, N, 1] → reshape 为 [B*T, N] pred_flat pred.view(-1, pred.size(2)) # [B*T, N] # 计算每个节点净流量流入 - 流出 net_flow torch.matmul(adj_matrix.T, pred_flat.T).T - torch.matmul(adj_matrix, pred_flat.T).T # 约束净流量接近 0理想状态无堆积 cons_loss torch.mean(torch.abs(net_flow)) return alpha * cons_loss F.l1_loss(pred, true)提示adj_matrix需根据实际路网构建。若你用 PeMSD7其W文件已提供加权邻接关系若用自有数据可用scikit-learn的NearestNeighbors按地理距离生成初始adj_matrix再人工校验主干道流向。4. 使用 PeMSD7 数据集跑通全流程——从解压到验证 RMSE 12.54.1 环境配置与依赖安装Linux/macOS本项目要求明确版本避免 PyTorch 与 CUDA 版本错配# 创建隔离环境 conda create -n sttransformer python3.8 conda activate sttransformer # 安装指定版本经实测兼容 pip install torch1.12.1cu113 torchvision0.13.1cu113 -f https://download.pytorch.org/whl/torch_stable.html pip install numpy1.21.6 pandas1.3.5 scikit-learn1.0.2 tqdm4.64.1 # 验证 CUDA 可用性 python -c import torch; print(torch.cuda.is_available(), torch.version.cuda) # 应输出 True 11.34.2 数据预处理3 步生成模型可读的.npz文件PeMSD7 原始数据为 CSV需转为压缩 NumPy 格式以加速 I/O# 进入 data/ 目录 cd data/ # 步骤1生成传感器 ID 映射确保顺序一致 python -c import pandas as pd df pd.read_csv(PeMSD7_V_228.csv, headerNone) print(Sensor count:, df.shape[1]) # 应输出 228 print(Time steps:, df.shape[0]) # 应输出 16992约 12 天 # 步骤2运行预处理脚本项目自带 preprocess.py python ../preprocess.py \ --data_dir ./ \ --output_dir ./processed/ \ --sensor_count 228 \ --input_len 12 \ --output_len 3 \ --train_ratio 0.6 \ --val_ratio 0.2 # 步骤3检查输出文件 ls ./processed/ # 应看到train.npz, val.npz, test.npz, scaler.pkl归一化参数preprocess.py关键逻辑读取PeMSD7_V_228.csv按列索引分配传感器 ID第 0 列 sensor_0对每列做 MinMaxScaler 归一化非 StandardScaler因流量非正态分布滑动窗口切片train.npz包含x_train输入和y_train标签两个数组。4.3 模型训练超参数调优的 4 个必调项train.py中以下参数直接影响收敛速度与最终 RMSE参数名默认值推荐调整范围作用说明--batch_size3216, 32, 64内存受限时选 16GPU 显存 ≥16GB 可试 64提升吞吐--learning_rate0.0010.0005, 0.001, 0.002初始用 0.001若 loss 振荡剧烈则降为 0.0005--weight_decay0.00010.00001, 0.0001防止过拟合交通数据噪声大不宜设过高--dropout0.10.1, 0.2, 0.3Spatial Block 对 dropout 更敏感超过 0.2 易欠拟合# 启动训练关键命令 python train.py \ --data_dir data/processed/ \ --model_path models/sttransformer_best.pth \ --epochs 100 \ --batch_size 32 \ --lr 0.001 \ --dropout 0.1 \ --save_every 10 \ --device cuda:0 # 训练日志关键指标第 100 轮应达到 # Train Loss: 0.0214 | Val MAE: 11.87 | Val RMSE: 12.434.4 测试与可视化用plot_results.py验证物理合理性训练完成后必须验证预测结果是否符合交通常识# 生成测试集预测图 python plot_results.py \ --model_path models/sttransformer_best.pth \ --data_dir data/processed/ \ --output_dir results/ \ --sensor_id 102 \ # 任选一个高频传感器 --num_plots 5 # 查看生成的 PNG ls results/ # 应看到sensor_102_pred_0.png, sensor_102_pred_1.png ...plot_results.py会绘制三组曲线蓝色实线真实流量ground truth橙色虚线模型预测值prediction灰色阴影区预测区间通过蒙特卡洛 Dropout 估计不确定性注意若在早高峰7:00–9:00时段预测曲线持续低于真实值 15%说明模型未充分学习周期性——需检查data_loader.py中是否漏掉了day_of_week特征编码。5. 将 ST-Transformer 迁移到自有数据——3 类常见路网的适配方案5.1 地磁/线圈检测器数据只需 2 处修改即可接入你的数据格式为sensor_id, timestamp, flow, speedCSV适配步骤重写data_loader.py中的load_dataset()函数替换原始 PeMSD7 读取逻辑def load_dataset(data_dir): # 读取所有传感器 CSV files glob.glob(f{data_dir}/sensor_*.csv) df_list [] for f in files: df pd.read_csv(f) df[sensor_id] int(f.split(_)[-1].split(.)[0]) # 从文件名提取 ID df_list.append(df) full_df pd.concat(df_list, ignore_indexTrue) # 按 sensor_id 和 timestamp 排序生成 (T, N) 矩阵 pivot_df full_df.pivot(indextimestamp, columnssensor_id, valuesflow) return pivot_df.fillna(0).values # [T, N]修改preprocess.py中的--sensor_count参数设为你实际传感器总数如 47重新运行预处理与训练无需改动模型结构。5.2 高精度浮动车数据GPS 轨迹必须增加 OD 矩阵构建GPS 数据含经纬度需先聚合为路段级流量# utils/od_builder.py 中核心逻辑 def build_od_matrix(traj_df, road_network_shp, time_bin30T): # traj_df: GPS 轨迹含 vehicle_id, lon, lat, timestamp # road_network_shp: 路网 Shapefile含路段 ID 与几何 # 步骤1将轨迹点匹配到最近路段用 GeoPandas.sjoin_nearest matched gpd.sjoin_nearest(traj_df, road_network_shp, distance_coldist) # 步骤2按时间窗口统计各路段进出量 od_matrix matched.groupby([timestamp_bin, road_id]).size().unstack(fill_value0) return od_matrix # [T, R]R 为路段数提示OD 矩阵维度R通常远大于传感器数N此时需在model.py中将num_sensors改为R并增大d_model至 128 以容纳更多空间信息。5.3 多模态数据融合如何加入天气与事件数据交通流受天气降雨量、能见度和事件施工、封路强影响。扩展特征维度# 修改 data_loader.py 中的 feature engineering def add_external_features(df): # df: [T, N] weather_df pd.read_csv(weather.csv, index_col0) # [T, 3]: rain, temp, visibility event_df pd.read_csv(events.csv, index_col0) # [T, 1]: is_construction # 拼接为 [T, N4]最后 4 列为外部特征 ext_features pd.concat([weather_df, event_df], axis1) return np.concatenate([df, np.tile(ext_features.values, (1, df.shape[1]))], axis1)此时模型输入维度变为(B, T, N, C4)需同步修改model.py中d_model的初始投影层# 在 STTransformerEncoder.__init__() 中 self.input_proj nn.Linear(C 4, d_model) # 原为 nn.Linear(C, d_model)最终验证时若加入天气特征后 RMSE 下降 8%说明模型成功捕获了气象对通行能力的影响——这是纯时序模型无法做到的物理可解释性突破。本文还有配套的精品资源点击获取
返回列表