Temporal Fusion Transformers (TFT) 中的 GRN 模块详解:输入、输出和 PyTorch 代码实现
Temporal Fusion Transformers (TFT) 中的 GRN 模块详解:输入、输出和 PyTorch 代码实现
在论文 'Temporal Fusion Transformers for Interpretable Multi-horizon Time Series Forecasting' 中定义的 TFT 架构中,GRN 模块扮演着重要的角色。它用于生成时间序列数据的邻接矩阵,为后续的注意力机制提供信息。
GRN 模块的输入和输出
GRN 模块的输入是一个形状为 (batch_size, num_nodes, embedding_size) 的张量,其中:
- batch_size: 批次数。
- num_nodes: 时间序列中的节点数。
- embedding_size: 节点的特征维度。
GRN 模块的输出是一个形状为 (batch_size, num_nodes, num_nodes) 的邻接矩阵,表示节点之间的关系。
PyTorch 代码实现
import torch
import torch.nn as nn
class GRN(nn.Module):
def __init__(self, embedding_size, num_nodes):
super(GRN, self).__init__()
self.fc1 = nn.Linear(embedding_size, num_nodes)
self.relu = nn.ReLU()
self.fc2 = nn.Linear(num_nodes, num_nodes)
def forward(self, x):
# x: (batch_size, num_nodes, embedding_size)
x = self.fc1(x) # (batch_size, num_nodes, num_nodes)
x = self.relu(x)
x = self.fc2(x) # (batch_size, num_nodes, num_nodes)
return x
在这个示例代码中,我们定义了一个 GRN 模块,它接受一个形状为 (batch_size, num_nodes, embedding_size) 的输入张量 x。在模块内部,我们首先使用一个全连接层将每个节点的特征嵌入到一个 num_nodes 维的向量中,然后使用 ReLU 激活函数进行非线性变换。最后,我们使用另一个全连接层将变换后的向量再次映射到 num_nodes 维,得到一个形状为 (batch_size, num_nodes, num_nodes) 的邻接矩阵。最终,我们返回这个邻接矩阵作为 GRN 模块的输出。
GRN 模块在 TFT 中的作用
需要注意的是,GRN 模块仅仅是一个邻接矩阵的生成器,它并不涉及到时间序列的处理。在 TFT 模型中,我们使用 GRN 模块来生成时间序列的邻接矩阵,然后将邻接矩阵作为第一个注意力层的输入,用于对时间序列的特征进行交互和整合。
通过这种方式,GRN 模块帮助 TFT 模型学习时间序列中各个节点之间的关系,并利用这些关系来进行更准确的多时间步预测。
原文地址: https://www.cveoy.top/t/topic/lRo0 著作权归作者所有。请勿转载和采集!