ARTICLE DETAIL

资讯详情

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

基于深度学习的时间序列预测:Temporal Fusion Transformer 实战配置与验证

基于深度学习的时间序列预测:Temporal Fusion Transformer 实战配置与验证 1. 从一次电力负荷预测翻车说起Temporal Fusion Transformer 到底解决什么问题去年帮一个做园区能耗管理的团队调模型他们的需求很具体用过去一周的用电数据预测未来 24 小时每小时的负荷而且要求给出预测区间不能只给一个点估计。他们一开始用的是 LSTM 加全连接层的组合单变量输入效果勉强能看但一遇到节假日或者气温骤降预测曲线就完全跑偏。更麻烦的是运维同事问“为什么这个时段预测偏高”没人能解释清楚模型就是个黑盒。这个场景其实非常典型。多变量时间序列预测的难点从来不只是“拟合一条曲线”而是要把三类信息揉在一起随时间变化且提前已知的变量比如小时、星期几、月份、节假日标记、随时间变化但预测时未知的变量比如实际用电量、实时气温、以及不随时间变化的静态变量比如每个电表对应的用户类型、区域编号。传统 ARIMA 类方法要求序列平稳处理外生变量很吃力普通 LSTM 虽然能吃多变量但对静态特征和已知未来特征的区分能力弱而且缺乏可解释性。Temporal Fusion TransformerTFT就是冲着这些痛点来的。它基于 Transformer 的自注意力机制专门为多步预测设计支持上面说的全部四类特征时变已知、时变未知、静态类别、静态实数。更关键的是它内置了变量选择网络和可解释的多头注意力训练完之后你能直接看到哪些特征重要、过去哪些时间步对当前预测影响大。对于需要向业务方解释预测依据的场景这一点比单纯刷低 MSE 有价值得多。我试过在同一个电力数据集上把 LSTM 和 TFT 做对比TFT 在加入静态的 consumer_id 和时变已知的 hour、day_of_week 之后验证集上的分位数损失明显更低而且预测曲线能跟上每日的周期性波动。下面我就把从数据窗口构造到训练回测的完整流程拆开讲配置片段可以直接复制。2. 前置准备TaoToken 接入与 PyTorch Forecasting 环境搭建在开始写模型之前先把两件事搞定一个是模型训练需要的 Python 环境另一个是如果你打算用 API 方式调用大模型辅助调试代码或生成配置需要一个稳定的接入点。这里我用的 TaoToken 来做模型对话和代码辅助它的 API 地址是 https://taotoken.net/api兼容 OpenAI 风格的请求格式配置起来比较直接。先说环境。TFT 的实现我推荐用 PyTorch Forecasting 这个库它把 TimeSeriesDataSet 和 TemporalFusionTransformer 都封装好了省去自己写 Dataset 的麻烦。安装命令如下注意版本对齐不然容易出兼容问题pip install torch2.0.1cu118 pytorch-lightning2.0.2 pytorch_forecasting1.0.0如果你没有 GPU把cu118去掉装 CPU 版本也能跑只是训练慢一些。装完之后验证一下import torch import pytorch_forecasting print(torch.__version__) print(pytorch_forecasting.__version__)接下来是 TaoToken 的接入。如果你只是本地跑模型这一步可以跳过但如果你想让大模型帮你检查 TimeSeriesDataSet 的参数配置、或者根据报错生成修复建议配好 API 会方便很多。在项目根目录建一个.env文件写入TAOTOKEN_API_KEY你的APIKey TAOTOKEN_BASE_URLhttps://taotoken.net/api然后在 Python 里这样调用import os from openai import OpenAI client OpenAI( api_keyos.getenv(TAOTOKEN_API_KEY), base_urlos.getenv(TAOTOKEN_BASE_URL) ) response client.chat.completions.create( modelgpt-4o, messages[ {role: user, content: TimeSeriesDataSet 的 max_encoder_length 和 max_prediction_length 分别代表什么} ] ) print(response.choices[0].message.content)API Key 的获取入口在 https://taotoken.net/api-keys登录后新建一个就行。模型对话的入口在 https://taotoken.net/models可以先用它跑通一个最小请求确认 Key 和 Base URL 没问题再去接 PyTorch Forecasting 的调试流程。这样分工的好处是模型训练本身不依赖网络但遇到配置报错时能快速拿到排查思路。环境这块还有一个坑PyTorch Forecasting 依赖 PyTorch Lightning而 Lightning 2.x 和 1.x 的 Trainer 参数有差异。上面锁定的 2.0.2 版本里acceleratorgpu和devices1是标准写法如果你装的是 1.9 以下的版本得改成gpus1。这个后面在训练配置里会再提。3. 可复制配置TimeSeriesDataSet 窗口构造与 TFT 模型参数这一节是核心我把数据窗口构造和模型初始化的配置完整写出来你可以直接改路径和列名套用到自己的数据上。先看数据格式。TFT 要求的数据是一个长表long format每一行是一个时间点必须有一个递增的time_idx列以及一个group_ids列来区分不同的时间序列。假设你的原始数据是宽表每个电表一列需要先 melt 成长表。下面是一个最小可运行的配置示例用 JSON 风格描述字段映射方便你对照自己的数据{ time_idx: hours_from_start, target: power_usage, group_ids: [consumer_id], static_categoricals: [consumer_id], time_varying_known_reals: [hours_from_start, day, day_of_week, month, hour], time_varying_unknown_reals: [power_usage], max_encoder_length: 168, max_prediction_length: 24, target_normalizer: GroupNormalizer }对应的 Python 构造代码from pytorch_forecasting import TimeSeriesDataSet from pytorch_forecasting.data import GroupNormalizer max_prediction_length 24 max_encoder_length 7 * 24 training_cutoff time_df[hours_from_start].max() - max_prediction_length training TimeSeriesDataSet( time_df[lambda x: x.hours_from_start training_cutoff], time_idxhours_from_start, targetpower_usage, group_ids[consumer_id], min_encoder_lengthmax_encoder_length // 2, max_encoder_lengthmax_encoder_length, min_prediction_length1, max_prediction_lengthmax_prediction_length, static_categoricals[consumer_id], time_varying_known_reals[hours_from_start, day, day_of_week, month, hour], time_varying_unknown_reals[power_usage], target_normalizerGroupNormalizer( groups[consumer_id], transformationsoftplus ), add_relative_time_idxTrue, add_target_scalesTrue, add_encoder_lengthTrue, ) validation TimeSeriesDataSet.from_dataset( training, time_df, predictTrue, stop_randomizationTrue ) batch_size 64 train_dataloader training.to_dataloader(trainTrue, batch_sizebatch_size, num_workers0) val_dataloader validation.to_dataloader(trainFalse, batch_sizebatch_size * 10, num_workers0)这里有几个参数值得展开。max_encoder_length168表示回看过去 168 小时也就是一周max_prediction_length24表示预测未来 24 小时。GroupNormalizer按consumer_id分组做归一化因为不同电表的用电量量级差异很大不归一化的话模型会被大数值的序列主导。transformationsoftplus保证归一化后的值非负适合用电量这种物理上不能为负的目标。模型初始化配置from pytorch_forecasting.models import TemporalFusionTransformer from pytorch_forecasting.metrics import QuantileLoss tft TemporalFusionTransformer.from_dataset( training, learning_rate0.001, hidden_size160, attention_head_size4, dropout0.1, hidden_continuous_size160, output_size7, lossQuantileLoss(), log_interval10, reduce_on_plateau_patience4, )output_size7对应 7 个分位数[0.02, 0.1, 0.25, 0.5, 0.75, 0.9, 0.98]这样模型输出的不只是点预测还有预测区间。hidden_size160和attention_head_size4跟原论文保持一致显存不够的话可以降到 64 和 2。训练用 PyTorch Lightning 的 Trainerimport pytorch_lightning as pl from pytorch_lightning.callbacks import EarlyStopping, LearningRateMonitor from pytorch_lightning.loggers import TensorBoardLogger early_stop_callback EarlyStopping( monitorval_loss, min_delta1e-4, patience5, verboseTrue, modemin ) lr_logger LearningRateMonitor() logger TensorBoardLogger(lightning_logs) trainer pl.Trainer( max_epochs45, acceleratorgpu, devices1, enable_model_summaryTrue, gradient_clip_val0.1, callbacks[lr_logger, early_stop_callback], loggerlogger, ) trainer.fit(tft, train_dataloaderstrain_dataloader, val_dataloadersval_dataloader)如果你用的是 CPU把acceleratorgpu改成acceleratorcpudevices1保留。gradient_clip_val0.1对 Transformer 类模型很重要能防止梯度爆炸。训练完成后保存最佳 checkpointbest_model_path trainer.checkpoint_callback.best_model_path best_tft TemporalFusionTransformer.load_from_checkpoint(best_model_path)这套配置我在 5 个电表、约 6000 小时的数据上跑过单卡 6 个 epoch 左右 EarlyStopping 就触发了验证损失稳定在 6.0 附近。下面讲怎么验证请求是否真的成功。4. 验证请求与成功结果预测、指标对比与可视化训练完不等于模型可用必须做验证。第一步是确认数据加载器吐出来的 batch 结构符合预期。在训练前可以先跑一个检查x, y next(iter(train_dataloader)) print(x[encoder_target].shape) print(x[decoder_target].shape) print(x[groups].shape)正常输出应该是encoder_target为[batch_size, encoder_length]decoder_target为[batch_size, prediction_length]。如果这里报错多半是time_idx不连续或者group_ids有缺失值。模型评估用分位数损失先算基准模型做对比。基准模型很简单直接用前一天的同一时段值作为预测。这个基准经常被忽略但在时间序列里它往往出奇地强import torch from pytorch_forecasting.metrics import Baseline actuals torch.cat([y[0] for x, y in iter(val_dataloader)]).to(cuda) baseline_predictions Baseline().predict(val_dataloader) baseline_loss (actuals - baseline_predictions).abs().mean().item() print(fBaseline MAE: {baseline_loss:.4f})然后算 TFT 的 P50 损失predictions best_tft.predict(val_dataloader) tft_loss (actuals - predictions).abs().mean().item() print(fTFT P50 MAE: {tft_loss:.4f})我实测下来基准模型 MAE 在 25 左右TFT 能降到 6 左右提升非常明显。如果你想看每个时间序列的单独损失per_series_loss (actuals - predictions).abs().mean(axis1) print(per_series_loss)输出是一个长度为 5 的张量对应 5 个电表。量级大的电表损失绝对值会高一些这是正常的可以再除以各自的平均功率做归一化对比。可视化部分用plot_prediction把预测区间和注意力权重一起画出来import matplotlib.pyplot as plt raw_predictions best_tft.predict(val_dataloader, moderaw, return_xTrue) for idx in range(5): fig, ax plt.subplots(figsize(10, 4)) best_tft.plot_prediction( raw_predictions.x, raw_predictions.output, idxidx, add_loss_to_titleQuantileLoss(), axax, ) plt.tight_layout() plt.savefig(fprediction_consumer_{idx}.png)图里灰色线是注意力分数能看出模型在预测某个时刻时过去哪些时间步的权重高。如果每日周期性明显你会看到每隔 24 小时出现一个小峰值。这一步跑通说明整个链路从数据窗口到预测输出都是通的。样本外预测稍微麻烦一点需要手动构造 decoder 数据。核心思路是取最后 168 小时作为 encoder 输入然后为未来 24 小时构造占位行已知特征hour、day_of_week 等填真实值未知特征power_usage填最后观测值import pandas as pd import numpy as np encoder_data time_df[lambda x: x.hours_from_start x.hours_from_start.max() - max_encoder_length] last_data time_df[lambda x: x.hours_from_start x.hours_from_start.max()] decoder_data pd.concat( [last_data.assign(datelambda x: x.date pd.offsets.Hour(i)) for i in range(1, max_prediction_length 1)], ignore_indexTrue, ) decoder_data[hours_from_start] ( (decoder_data[date] - earliest_time).dt.seconds / 3600 (decoder_data[date] - earliest_time).dt.days * 24 ).astype(int) decoder_data[hours_from_start] encoder_data[hours_from_start].max() 1 - decoder_data[hours_from_start].min() decoder_data[month] decoder_data[date].dt.month.astype(np.int64) decoder_data[hour] decoder_data[date].dt.hour.astype(np.int64) decoder_data[day] decoder_data[date].dt.day.astype(np.int64) decoder_data[day_of_week] decoder_data[date].dt.dayofweek.astype(np.int64) new_prediction_data pd.concat([encoder_data, decoder_data], ignore_indexTrue) new_prediction_data new_prediction_data.query(consumer_id MT_002) new_raw_predictions best_tft.predict(new_prediction_data, moderaw, return_xTrue) best_tft.plot_prediction( new_raw_predictions.x, new_raw_predictions.output, idx0, show_future_observedFalse, )跑完这一步你会得到一张未来 24 小时的预测曲线带 7 个分位数的区间。如果曲线平滑且区间合理说明模型没有过拟合到训练集的噪声上。5. 本篇常见错排查401、local proxy failed、reading choices 与 OAuth 报错这一节把我踩过的坑列出来对照报错信息找解决方案。401 Unauthorized如果你在调用 TaoToken API 做代码辅助时遇到这个先检查.env里的TAOTOKEN_API_KEY有没有多余空格再确认base_url写的是https://taotoken.net/api而不是带其他路径。用 curl 快速验证curl https://taotoken.net/api/models \ -H Authorization: Bearer $TAOTOKEN_API_KEY如果返回模型列表说明 Key 没问题如果还是 401去 https://taotoken.net/api-keys 重新生成一个。local proxy failed这个报错通常出现在你本地开了某些网络工具导致请求发不出去。解决方式是检查环境变量里有没有HTTP_PROXY或HTTPS_PROXY临时清掉unset HTTP_PROXY unset HTTPS_PROXY然后在 Python 里显式指定不走代理import os os.environ[NO_PROXY] taotoken.netreading choices 报错这个一般出现在解析 API 返回时response.choices为空。原因可能是请求被截断或者模型名写错。检查model参数是否拼写正确比如gpt-4o不要写成gpt4o。另外确认max_tokens没有设成 0。OAuth 相关报错如果你在用某些 CLI 工具比如 Claude Code 或 Codex 的认证流程时遇到 OAuth 失败先确认本地时间是否准确OAuth token 对时间偏差敏感。然后检查配置文件路径是否正确。以 Codex 的auth.json为例它通常放在~/.config/codex/auth.json内容格式{ base_url: https://taotoken.net/api, api_key: 你的APIKey, model: gpt-4o }三件套缺一不可Base URL、Key、Model ID。如果你用的是 Cline 或 CC Switch 这类工具MCP 配置里同样要写全这三项。少写 Model ID 会导致请求发出去但返回空。TimeSeriesDataSet 报错 “time_idx must be consecutive”这个不是 API 问题是数据问题。检查每个group_id下的time_idx是否从 0 开始且步长为 1。如果有缺失小时需要补行或者调整time_idx。CUDA out of memory把batch_size从 64 降到 32 或 16同时把hidden_size从 160 降到 64。如果还是不够用acceleratorcpu先跑通流程。验证损失不下降先检查GroupNormalizer有没有加再确认learning_rate是不是太大。TFT 对学习率比较敏感0.001 是安全值0.01 以上容易震荡。6. 从训练到回测的完整链路与后续调优方向把上面几步串起来一个可复现的 TFT 实战流程就成型了数据预处理成长表、TimeSeriesDataSet 构造窗口、TFT 初始化与训练、基准对比与可视化验证、样本外预测。我在电力负荷数据上跑完这套流程从原始 CSV 到出预测图大概半天时间其中大部分花在数据清洗和特征构造上模型训练本身很快。如果你想让效果再上一个台阶有几个方向可以试。一是超参数搜索PyTorch Forecasting 内置了 Optuna 集成from pytorch_forecasting.models.temporal_fusion_transformer.tuning import optimize_hyperparameters study optimize_hyperparameters( train_dataloader, val_dataloader, model_pathoptuna_test, n_trials20, max_epochs10, gradient_clip_val_range(0.01, 1.0), hidden_size_range(30, 128), hidden_continuous_size_range(30, 128), attention_head_size_range(1, 4), learning_rate_range(0.001, 0.1), dropout_range(0.1, 0.3), reduce_on_plateau_patience4, use_learning_rate_finderFalse, ) print(study.best_trial.params)注意这个很吃显存n_trials别设太大先跑 5 到 10 次看看趋势。二是特征工程把节假日标记、气温、电价等外生变量加进time_varying_known_realsTFT 的变量选择网络会自动评估它们的重要性。三是用interpret_output做特征重要性分析interpretation best_tft.interpret_output(raw_predictions.output, reductionsum) best_tft.plot_interpretation(interpretation)这张图能直接告诉你哪些特征对预测贡献大。如果consumer_id的重要性很低说明你的多个时间序列其实可以用一个全局模型建模不需要单独训练。如果某个时变已知特征重要性异常高检查一下它是不是泄露了未来信息。最后提醒一点TFT 虽然强但不是所有场景都需要它。如果你的数据只有单变量、没有外生特征、序列也很平稳ARIMA 或者简单的指数平滑可能更快更稳。TFT 的价值在于多变量、多步、需要可解释性的复杂场景。选型之前先用基准模型跑一遍确认简单方法不够用再上 TFT这样投入产出比最合理。
返回列表