正在加载中...

展开本页目录
算法教程TFT-Temporal Fusion Transformer

TFT-Temporal Fusion Transformer

No.158 · 在线教程

当前项目中的 TFT-Temporal Fusion Transformer,并不是论文原版那种带有静态变量编码、变量选择网络、门控残差网络、可解释多头注意力、分位数多步输出的完整 TFT。真实代码实现位于 core/tftcalculator.py,其结构更接近一个轻量级时序…

TFT-Temporal Fusion Transformer

1. 方法概述

当前项目中的 TFT-Temporal Fusion Transformer,并不是论文原版那种带有静态变量编码、变量选择网络、门控残差网络、可解释多头注意力、分位数多步输出的完整 TFT。真实代码实现位于 core/tft_calculator.py,其结构更接近一个轻量级时序注意力模型

  • 数值特征滑动窗口输入;
  • 线性层做特征嵌入;
  • 单层 LSTM 编码序列;
  • 以最后一个时间步隐藏状态作为 query,对整段历史做一次 MultiheadAttention
  • 最后接全连接层输出回归值或分类 logits。

设原始数据按时间顺序为

$$ \mathcal{D}=\{(x_t,y_t)\}_{t=1}^{N},\qquad x_t\in\mathbb{R}^{d},\ y_t\in\mathcal{Y} \tag{1} $$

其中 \(x_t\) 是第 \(t\) 行的特征向量,\(y_t\) 是目标变量。界面上虽然名为 Temporal Fusion Transformer,但实现层面没有标准 TFT 中的“fusion”模块化结构,因此写论文时更准确的称呼应是:TFT 风格的轻量时序注意力网络

2. 数据处理与窗口构造

2.1 特征筛选

程序先删除目标列,仅保留剩余列中的数值特征,记为

$$ \tilde{x}_t\in\mathbb{R}^{d'} \tag{2} $$

非数值特征会被静默丢弃,不会做 One-Hot 编码,也不会构造静态协变量或已知未来协变量。如果去掉目标列后不存在数值特征,则直接报错。

2.2 滑动窗口

设输入窗口长度为 \(L=\text{seq\_len}\),预测步长为 \(P=\text{pred\_len}\)。代码对每个可行位置 \(i\) 构造

$$ \mathbf{X}_i= \begin{bmatrix} \tilde{x}_{i-L}\\ \tilde{x}_{i-L+1}\\ \vdots\\ \tilde{x}_{i-1} \end{bmatrix} \in\mathbb{R}^{L\times d'} \tag{3} $$

对应标签取为

$$ z_i=y_{i+P-1} \tag{4} $$

因此总样本数为

$$ M=N-L-P+1 \tag{5} $$

这说明 pred_len 在当前实现里并不是“输出未来 \(P\) 步序列”,而是“选未来第 \(P\) 步那个单点作为监督目标”,整体仍是 sequence-to-one 结构。

2.3 分类标签编码

分类任务中,程序使用 LabelEncoder 将标签编码为

$$ \phi:\mathcal{Y}\to\{0,1,\ldots,C-1\} \tag{6} $$

并把映射关系保存到 class_mapping。测试预测表在分类任务下还会额外导出 actual_labelpredicted_label 两列。

2.4 全量窗口标准化

若勾选 normalize=True,程序会在全部窗口构造完成后、切分训练测试集之前,对特征做标准化:

$$ \mu=\frac{1}{ML}\sum_{i=1}^{M}\sum_{\tau=1}^{L}\mathbf{X}_{i,\tau} \tag{7} $$

$$ \sigma=\sqrt{\frac{1}{ML}\sum_{i=1}^{M}\sum_{\tau=1}^{L}\left(\mathbf{X}_{i,\tau}-\mu\right)^2} \tag{8} $$

$$ \mathbf{X}_i^{\ast}=\frac{\mathbf{X}_i-\mu}{\sigma} \tag{9} $$

其中零方差维度被替换为 1。由于 \(\mu,\sigma\) 使用了全量窗口样本,因此这里存在信息泄漏。

2.5 随机切分而非时序留出

标准化完成后,程序通过 train_test_split(..., shuffle=True) 进行训练/测试切分:

$$ (\mathcal{X}_{\mathrm{tr}},\mathcal{X}_{\mathrm{te}},\mathcal{y}_{\mathrm{tr}},\mathcal{y}_{\mathrm{te}}) =\mathrm{Split}(\{\mathbf{X}_i^{\ast},z_i\}_{i=1}^{M}) \tag{10} $$

分类任务若类别数大于 1,则会启用 stratify=y。这意味着该模块虽然输入是时序窗口,但评估并不是严格意义上的“按时间向前外推测试”,而是随机打散后的监督学习式评估。

3. 轻量化 TFT 风格网络

3.1 线性嵌入

输入窗口中每个时间步的特征先经过线性映射:

$$ e_{i,\tau}=W_e x_{i,\tau}^{\ast}+b_e,\qquad e_{i,\tau}\in\mathbb{R}^{h} \tag{11} $$

其中 \(h=\text{hidden\_size}\)。代码要求

$$ h \bmod H = 0 \tag{12} $$

这里 \(H=\text{n\_heads}\) 是注意力头数,因为 MultiheadAttention 要求隐藏维度能被头数整除。

3.2 LSTM 编码

嵌入序列送入单层 LSTM

$$ \mathbf{h}_{i,1:L}=\mathrm{LSTM}(e_{i,1:L}) \tag{13} $$

得到每个时间步的隐藏表示 \(\mathbf{h}_{i,\tau}\in\mathbb{R}^{h}\)。

3.3 最后一步查询的多头注意力

模型并没有 TFT 原论文中的解码器与可解释注意力模块,而是简单取最后一个时间步作为 query:

$$ q_i=\mathbf{h}_{i,L} \tag{14} $$

将整段隐藏序列同时作为 key 和 value,做一次注意力汇聚:

$$ \alpha_{i,\tau}=\mathrm{softmax}\!\left(\frac{q_i^\top W_Q^\top W_K \mathbf{h}_{i,\tau}}{\sqrt{h/H}}\right) \tag{15} $$

$$ c_i=\sum_{\tau=1}^{L}\alpha_{i,\tau} W_V \mathbf{h}_{i,\tau} \tag{16} $$

实际代码调用的是 nn.MultiheadAttention(query, h, h),其中 query 的长度只有 1,因此整个模型本质上是在做“最后一步状态对历史序列的一次注意力读出”。

3.4 输出层

注意力输出先经过 dropout,再送入线性层:

$$ \hat{y}_i=W_o\,\mathrm{Dropout}(c_i)+b_o \tag{17} $$

回归任务中 \(\hat{y}_i\in\mathbb{R}\);分类任务中 \(\hat{y}_i\in\mathbb{R}^{C}\),通过 argmax 生成预测类别。

4. 与标准 TFT 的主要差异

从源码角度看,当前实现缺少标准 TFT 中的以下关键组成:

  • 静态协变量编码;
  • 过去已知/未来已知变量分组;
  • 变量选择网络(Variable Selection Network);
  • 门控残差网络(GRN)与门控跳连;
  • 多步解码器;
  • 分位数预测头;
  • 原论文中的可解释注意力分析输出。

因此它虽然保留了 “LSTM + Attention” 的组合风格,但并不是论文式的完整 TFT 实现。

5. 训练目标、评价指标与输出结果

5.1 训练目标

回归任务使用均方误差损失:

$$ \mathcal{L}_{\mathrm{reg}}=\frac{1}{|\mathcal{X}_{\mathrm{tr}}|}\sum_i(z_i-\hat{y}_i)^2 \tag{18} $$

分类任务使用交叉熵损失:

$$ \mathcal{L}_{\mathrm{cls}}=-\frac{1}{|\mathcal{X}_{\mathrm{tr}}|}\sum_i\log\frac{\exp(\hat{y}_{i,z_i})}{\sum_{c=1}^{C}\exp(\hat{y}_{i,c})} \tag{19} $$

训练过程中只记录 train_loss,没有单独验证集、没有早停,也没有学习率调度。

5.2 回归指标

回归任务输出 rmsemaer2

$$ \mathrm{RMSE}=\sqrt{\frac{1}{|\mathcal{X}_{\mathrm{te}}|}\sum_i(y_i-\hat{y}_i)^2} \tag{20} $$

$$ \mathrm{MAE}=\frac{1}{|\mathcal{X}_{\mathrm{te}}|}\sum_i|y_i-\hat{y}_i| \tag{21} $$

$$ R^2=1-\frac{\sum_i(y_i-\hat{y}_i)^2}{\sum_i(y_i-\bar{y})^2} \tag{22} $$

5.3 分类指标

分类任务输出 accuracyf1_macrof1_weighted。其中准确率为

$$ \mathrm{Accuracy}=\frac{1}{|\mathcal{X}_{\mathrm{te}}|}\sum_i\mathbf{1}(\hat{y}_i=y_i) \tag{23} $$

加权 F1 为

$$ \mathrm{F1}_{\mathrm{weighted}}=\sum_{c=1}^{C}\frac{n_c}{\sum_j n_j}\cdot \mathrm{F1}_c \tag{24} $$

5.4 输出结果说明

当前模块导出的 Excel 工作簿包含:

  • RawData
  • TestPreds
  • Metrics
  • Parameters
  • WindowedMeta
  • TrainingLog
  • Charts

其中:

  • 回归任务通常导出 pred_vs_true.pngresidual_hist.png
  • 分类任务会生成 confusion_matrix 图;
  • TrainingLog 只含 epochtrain_loss
  • WindowedMeta 会记录窗口数、特征列、是否标准化以及分类映射。

6. 输出结果与复现

结果页计算完成后,会自动生成 repro_*.py 脚本。该脚本会:

  1. 优先把原始输入文件复制到结果目录下的 repro_inputs/
  2. 若没有原始文件路径,则把 raw_data 另存为 CSV;
  3. 重新读取该输入文件;
  4. 调用同一个 TFTCalculator.run()save_to_excel() 生成 *_repro.xlsx

所以,这个复现脚本并不是“保存训练好的模型权重再加载推理”,而是按当前源码和当前参数重新训练一遍

7. 实现说明与注意事项

结合 core/tft_calculator.pyui/upload_widget.pyui/results_widget.py 和实际导出的结果文件,可以把该模块总结为:

  • 它确实使用了 LSTM + MultiheadAttention,但不是标准论文版 TFT;
  • pred_len 只决定未来第几步单点标签,不做多步联合输出;
  • 仅数值特征参与建模,非数值特征会被直接忽略;
  • 标准化先于训练/测试切分,存在信息泄漏;
  • 训练/测试切分使用随机打散,而非时间顺序留出;
  • 训练日志与复现脚本都较完整,但复现实质上是重新训练。

因此,在论文说明中更恰当的表述应是:该软件实现了一个带 LSTM 编码和注意力汇聚的 TFT 风格轻量模型,而不是原始 Temporal Fusion Transformer 的完整工程复现。

8. 论文写作模板

8.1 方法描述模板

可在论文“方法部分”中写为:

“本文采用轻量化 TFT 风格网络对时序样本进行预测或分类。首先,将原始序列样本按时间顺序构造成监督学习窗口,并完成训练、验证和测试划分;其次,利用 LSTM 编码与注意力汇聚结构对时序特征进行建模,并输出目标变量的预测结果;随后,根据任务类型采用相应损失函数完成模型训练;最后,在测试集上输出预测结果、评价指标和图表,用于分析模型的时序建模效果。”

8.2 结果解释模板

结果部分可写为:轻量化 TFT 通过 LSTM 编码与注意力汇聚兼顾局部动态与关键时间步信息。若注意力结构带来的性能提升有限,应在正文中说明当前实现与标准 TFT 存在明显简化差异,不宜直接做标准 TFT 的严格对照结论。

8.3 表格标题模板

表题可写为:轻量化 TFT 模型预测结果与性能指标汇总表。

8.4 图表题注模板

图注可写为:轻量化 TFT 模型真实值与预测值对比曲线。

8.5 表格示例

建议列名:时间点、真实值、预测值、残差、RMSE、MAE、\(R^2\) 或分类评价指标。

9. 论文写作建议

论文中建议把该模块写成“轻量化 TFT 风格模型”,并明确指出它不是原始 TFT 的完整复现。结果部分可重点展示注意力汇聚后的预测性能、测试集指标和预测图。若论文强调模型创新性,应避免直接把当前实现写成完整的 Temporal Fusion Transformer。

10. 单篇终审补充

10.1 图题与表题对齐建议

  • RawData 表可写为:表X 轻量化 TFT 原始数据表。
  • TestPreds 表可写为:表X 轻量化 TFT 测试集真实值与预测值对照。
  • Metrics 表可写为:表X 轻量化 TFT 性能指标汇总。
  • Parameters 表可写为:表X 轻量化 TFT 参数设置与运行配置。
  • WindowedMeta 表可写为:表X 轻量化 TFT 窗口构造与元信息记录。
  • TrainingLog 表可写为:表X 轻量化 TFT 训练损失记录。
  • Charts 表可写为:表X 轻量化 TFT 图表索引与路径清单。
  • pred_vs_true.png 建议写为:图X 轻量化 TFT 真实值与预测值对比图。
  • residual_hist.png 建议写为:图X 轻量化 TFT 残差分布图。

10.2 终审说明

  • 当前代表性结果目录中的真实主工作簿为 tft_baseline.xlsx,真实工作表为 RawData/TestPreds/Metrics/Parameters/WindowedMeta/TrainingLog/Charts,并不存在额外的注意力权重表、变量选择表或标准 TFT 专属解释表。
  • 当前代表性结果目录可见的实体图文件主要是 pred_vs_true.pngresidual_hist.png。正文若按分类任务写,还需确认是否实际生成 confusion_matrix 图,不能把分类图和回归图混写成同一批结果。
  • 当前目录中未直接看到 repro 脚本与 repro_inputs 副本,因此论文附录若需要展示复现实验,应优先引用同算法其他结果目录中的真实 repro_*.py 结构,而不是臆造当前目录已含 repro 文件。
  • TrainingLog 在当前实现中只记录 epochtrain_loss,不能被写成完整的深度学习训练监控面板;正文若讨论训练过程,应避免虚构验证损失曲线或早停细节。
  • 该模块本质是“TFT 风格轻量模型”,不是原始论文版 Temporal Fusion Transformer。终稿里若继续使用 TFT 名称,应在方法部分保留“轻量化 / 简化版 / 工程化实现”的限定语。

10.3 全量强化补充

本篇终审补充绑定的真实算法目录为 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer。本次采用两类真实结果证据:baseline 目录 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/__enhanced_tft_baseline_20260310,以及 runtime/repro 目录 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/_tft_enhanced_runtime_20260310/TFT-Temporal Fusion Transformer分析结果_20260318_232204

baseline 目录中的主结果工作簿为:

  • tft_baseline.xlsx

runtime/repro 目录中的工作簿为:

  • tft_results_20260318_232204.xlsx

两者实测工作表一致,均包含:

  • RawData
  • TestPreds
  • Metrics
  • Parameters
  • WindowedMeta
  • TrainingLog
  • Charts

这里需要特别说明:baseline 目录只有 tft_baseline.xlsxtft_baseline_plots/ 图目录;它本身没有 repro_*.py。复现实物位于另一个 runtime 目录中。因此这篇文档应明确拆分“baseline 结构证据”和“runtime/repro 结构证据”,不能把两者合并成同一个目录。

baseline 目录中的真实图文件为:

  • tft_baseline_plots/pred_vs_true.png
  • tft_baseline_plots/residual_hist.png

runtime/repro 目录中的真实图文件为:

  • tft_results_20260318_232204_plots/pred_vs_true.png
  • tft_results_20260318_232204_plots/residual_hist.png

两套图图义一致,都是预测对比图与残差直方图;区别在于所属运行目录,不是模型输出了更多图型。

复现实物方面,runtime/repro 目录实际包含:

  • 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/_tft_enhanced_runtime_20260310/TFT-Temporal Fusion Transformer分析结果_20260318_232204/repro_TFT-Temporal_Fusion_Transformer_20260318_232204.py
  • 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/_tft_enhanced_runtime_20260310/TFT-Temporal Fusion Transformer分析结果_20260318_232204/repro_inputs/tft_enhanced_input.csv

因此这一篇可以写成“已有独立 runtime/repro 目录验证相对输入副本”,但不能说 baseline 目录本身已经包含 repro 脚本。

11. 软件实现核查补充(2026-07)

  • 当前基线主结果目录应写作 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/manual_tft_baseline_20260318,runtime/repro 目录应写作 具体的算法3/深度学习与时序网络/TFT-Temporal Fusion Transformer/results/_tft_enhanced_runtime_20260310/TFT-Temporal Fusion Transformer分析结果_20260318_232204
  • 正文应围绕 RawDataProcessedTestPredsWindowedMetaTrainingLogCharts 来写。
  • 图证应对应 pred_vs_true.pngresidual_hist.png,并把 baseline 与 runtime/repro 目录分开说明。
  • 复现脚本应按 repro_TFT-Temporal_Fusion_Transformer_20260318_232204.py + repro_inputs/tft_enhanced_input.csv 的口径说明。