☰
基于深度学习的时间序列预测:Temporal Fusion Transformer 实战配置与验证
2026/10/1 6:52:25 网站建设 项目流程

1. 从一次电力负荷预测翻车说起:Temporal Fusion Transformer 到底解决什么问题

去年帮一个做园区能耗管理的团队调模型,他们的需求很具体:用过去一周的用电数据,预测未来 24 小时每小时的负荷,而且要求给出预测区间,不能只给一个点估计。他们一开始用的是 LSTM 加全连接层的组合,单变量输入,效果勉强能看,但一遇到节假日或者气温骤降,预测曲线就完全跑偏。更麻烦的是,运维同事问“为什么这个时段预测偏高”,没人能解释清楚,模型就是个黑盒。

这个场景其实非常典型。多变量时间序列预测的难点从来不只是“拟合一条曲线”,而是要把三类信息揉在一起:随时间变化且提前已知的变量(比如小时、星期几、月份、节假日标记)、随时间变化但预测时未知的变量(比如实际用电量、实时气温)、以及不随时间变化的静态变量(比如每个电表对应的用户类型、区域编号)。传统 ARIMA 类方法要求序列平稳,处理外生变量很吃力;普通 LSTM 虽然能吃多变量,但对静态特征和已知未来特征的区分能力弱,而且缺乏可解释性。

Temporal Fusion Transformer(TFT)就是冲着这些痛点来的。它基于 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 torch==2.0.1+cu118 pytorch-lightning==2.0.2 pytorch_forecasting==1.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_URL=https://taotoken.net/api

然后在 Python 里这样调用:

import os from openai import OpenAI client = OpenAI( api_key=os.getenv("TAOTOKEN_API_KEY"), base_url=os.getenv("TAOTOKEN_BASE_URL") ) response = client.chat.completions.create( model="gpt-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 版本里,accelerator='gpu'和devices=1是标准写法,如果你装的是 1.9 以下的版本,得改成gpus=1。这个后面在训练配置里会再提。

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_idx="hours_from_start", target="power_usage", group_ids=["consumer_id"], min_encoder_length=max_encoder_length // 2, max_encoder_length=max_encoder_length, min_prediction_length=1, max_prediction_length=max_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_normalizer=GroupNormalizer( groups=["consumer_id"], transformation="softplus" ), add_relative_time_idx=True, add_target_scales=True, add_encoder_length=True, ) validation = TimeSeriesDataSet.from_dataset( training, time_df, predict=True, stop_randomization=True ) batch_size = 64 train_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=0) val_dataloader = validation.to_dataloader(train=False, batch_size=batch_size * 10, num_workers=0)

这里有几个参数值得展开。max_encoder_length=168表示回看过去 168 小时,也就是一周;max_prediction_length=24表示预测未来 24 小时。GroupNormalizer按consumer_id分组做归一化,因为不同电表的用电量量级差异很大,不归一化的话模型会被大数值的序列主导。transformation="softplus"保证归一化后的值非负,适合用电量这种物理上不能为负的目标。

模型初始化配置:

from pytorch_forecasting.models import TemporalFusionTransformer from pytorch_forecasting.metrics import QuantileLoss tft = TemporalFusionTransformer.from_dataset( training, learning_rate=0.001, hidden_size=160, attention_head_size=4, dropout=0.1, hidden_continuous_size=160, output_size=7, loss=QuantileLoss(), log_interval=10, reduce_on_plateau_patience=4, )

output_size=7对应 7 个分位数[0.02, 0.1, 0.25, 0.5, 0.75, 0.9, 0.98],这样模型输出的不只是点预测,还有预测区间。hidden_size=160和attention_head_size=4跟原论文保持一致,显存不够的话可以降到 64 和 2。训练用 PyTorch Lightning 的 Trainer:

import pytorch_lightning as pl from pytorch_lightning.callbacks import EarlyStopping, LearningRateMonitor from pytorch_lightning.loggers import TensorBoardLogger early_stop_callback = EarlyStopping( monitor="val_loss", min_delta=1e-4, patience=5, verbose=True, mode="min" ) lr_logger = LearningRateMonitor() logger = TensorBoardLogger("lightning_logs") trainer = pl.Trainer( max_epochs=45, accelerator="gpu", devices=1, enable_model_summary=True, gradient_clip_val=0.1, callbacks=[lr_logger, early_stop_callback], logger=logger, ) trainer.fit(tft, train_dataloaders=train_dataloader, val_dataloaders=val_dataloader)

如果你用的是 CPU,把accelerator="gpu"改成accelerator="cpu",devices=1保留。gradient_clip_val=0.1对 Transformer 类模型很重要,能防止梯度爆炸。训练完成后保存最佳 checkpoint:

best_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(f"Baseline MAE: {baseline_loss:.4f}")

然后算 TFT 的 P50 损失:

predictions = best_tft.predict(val_dataloader) tft_loss = (actuals - predictions).abs().mean().item() print(f"TFT P50 MAE: {tft_loss:.4f}")

我实测下来,基准模型 MAE 在 25 左右,TFT 能降到 6 左右,提升非常明显。如果你想看每个时间序列的单独损失:

per_series_loss = (actuals - predictions).abs().mean(axis=1) print(per_series_loss)

输出是一个长度为 5 的张量,对应 5 个电表。量级大的电表损失绝对值会高一些,这是正常的,可以再除以各自的平均功率做归一化对比。

可视化部分,用plot_prediction把预测区间和注意力权重一起画出来:

import matplotlib.pyplot as plt raw_predictions = best_tft.predict(val_dataloader, mode="raw", return_x=True) for idx in range(5): fig, ax = plt.subplots(figsize=(10, 4)) best_tft.plot_prediction( raw_predictions.x, raw_predictions.output, idx=idx, add_loss_to_title=QuantileLoss(), ax=ax, ) plt.tight_layout() plt.savefig(f"prediction_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(date=lambda x: x.date + pd.offsets.Hour(i)) for i in range(1, max_prediction_length + 1)], ignore_index=True, ) 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_index=True) new_prediction_data = new_prediction_data.query("consumer_id == 'MT_002'") new_raw_predictions = best_tft.predict(new_prediction_data, mode="raw", return_x=True) best_tft.plot_prediction( new_raw_predictions.x, new_raw_predictions.output, idx=0, show_future_observed=False, )

跑完这一步,你会得到一张未来 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.net"

reading 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。如果还是不够,用accelerator="cpu"先跑通流程。

验证损失不下降:先检查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_path="optuna_test", n_trials=20, max_epochs=10, 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_patience=4, use_learning_rate_finder=False, ) print(study.best_trial.params)

注意这个很吃显存,n_trials别设太大,先跑 5 到 10 次看看趋势。二是特征工程,把节假日标记、气温、电价等外生变量加进time_varying_known_reals,TFT 的变量选择网络会自动评估它们的重要性。三是用interpret_output做特征重要性分析:

interpretation = best_tft.interpret_output(raw_predictions.output, reduction="sum") best_tft.plot_interpretation(interpretation)

这张图能直接告诉你哪些特征对预测贡献大。如果consumer_id的重要性很低,说明你的多个时间序列其实可以用一个全局模型建模,不需要单独训练。如果某个时变已知特征重要性异常高,检查一下它是不是泄露了未来信息。

最后提醒一点:TFT 虽然强,但不是所有场景都需要它。如果你的数据只有单变量、没有外生特征、序列也很平稳,ARIMA 或者简单的指数平滑可能更快更稳。TFT 的价值在于多变量、多步、需要可解释性的复杂场景。选型之前先用基准模型跑一遍,确认简单方法不够用,再上 TFT,这样投入产出比最合理。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询