Python技术迷

使用 TiDE 进行时间序列预测

时间序列这东西,大家应该都踩过坑:业务同学问一句“下个月的日活/订单量/电量能不能提前给个预测?”,你一看数据,季节性、节假日、活动全混在一起,用个 ARIMA 顶不住,用个 LSTM / Transformer 又又又慢还不好调。

最近两年挺火的一个模型叫 TiDE(Time-series Dense Encoder),就是专门干“长序列时间序列预测”的,而且结构极其接地气:全是 MLP,没有注意力,没有循环。这是 Google 团队 2023 年发的工作,在一堆公开数据集上能打过很多 Transformer,同时还快好几倍。

下面我就按一个“从业务到代码”的顺序,聊聊 TiDE 到底在干嘛,以及怎么用 Python 写个简化版跑起来。

你可以把 TiDE 想成一个带“记忆”的高配版线性回归:

  • 输入:一段过去的时间序列(比如过去 96 小时的用电),再加上各种特征:

    • 历史协变量:小时、星期几、是否周末、天气之类;
    • 未来协变量:未来 24 小时的这些特征(这个通常是“已知的未来”,比如节日日历)。
  • 输出:未来一段时间的预测(比如未来 24 小时的用电)。

TiDE 的关键点有三个:

  1. Encoder:用一堆带残差的 MLP,把“过去的 y + 过去的特征”压成一个高维向量,就相当于一个“历史总结”。
  2. Decoder:再用几层 MLP,把这个向量展开成“未来整体的一个隐向量”。
  3. Temporal Decoder:这一步比较骚,它会在每一个未来时间点上,把刚才那个隐向量和“该时间点的未来特征”拼在一起,再过 MLP 输出单步预测。这样每个时间点都能感知自己的“未来特征”,比如今天是节假日,就会单点往上抬。

同时,TiDE 还保留了一条线性残差通路:用一个线性层直接把历史 y 映射到未来 y,再叠加到最终输出上,相当于保证“最简单的线性模型”始终包含在 TiDE 里。

和传统模型比,它厉害在哪儿?

简单感受一下它解决的几个痛点:

  • 长序列友好:全是 MLP,计算复杂度基本跟序列长度线性相关,不会像自注意力那样 O(L²) 起飞。
  • 天然支持协变量:不论是历史特征还是未来特征,都能对到每个时间点上融合进去,尤其适合带节假日/促销/天气的业务场景。
  • 实现简单:没有 RNN、没有注意力,从代码层面就是堆 Linear + ReLU + Dropout + LayerNorm 的残差块,PyTorch 两下就写完。
  • 有现成库:像 Darts 的 TiDEModel、PyTorch Forecasting 的 TiDEModel、Nixtla 的 neuralforecast.models.TiDE 都已经把细节封好了,业务里直接上就行。

用 Python 写个“迷你 TiDE”

下面这个是极简版 TiDE,不跟论文一模一样,但思路是完全一样:

  • Encoder:若干残差块,把历史 flatten 成一个向量;
  • Decoder:把这个向量映射成“未来隐表示”;
  • Temporal decoder:融合未来协变量,一步步吐出未来的 y。

先是一个残差块:

import torch
import torch.nn as nn
import torch.nn.functional as F


classResidualBlock(nn.Module):
def__init__(self, in_dim, hidden_dim, dropout=0.1):
        super().__init__()
        self.fc1 = nn.Linear(in_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, in_dim)
        self.dropout = nn.Dropout(dropout)
        self.norm = nn.LayerNorm(in_dim)

defforward(self, x):
# x: (batch, dim)
        residual = x
        out = F.relu(self.fc1(x))
        out = self.dropout(out)
        out = self.fc2(out)
        out = self.dropout(out)
        out = self.norm(out + residual)
return out

然后是一个简化版 TiDE 模型,假设:

  • 过去窗口长度 input_len

  • 预测长度 output_len

  • 目标序列是单变量(比如某个城市的负载),但可以带:

    • 历史协变量 hist_cov_dim
    • 未来协变量 fut_cov_dim
classMiniTiDE(nn.Module):
def__init__(
        self,
        input_len: int,
        output_len: int,
        hist_cov_dim: int = 0,
        fut_cov_dim: int = 0,
        hidden_size: int = 128,
        enc_layers: int = 2,
        dec_layers: int = 2,
        temporal_hidden: int = 64,
    )
:

        super().__init__()
        self.input_len = input_len
        self.output_len = output_len
        self.hist_feat_dim = 1 + hist_cov_dim  # 1 for target y

# 编码输入总维度 = input_len * 每个时间点特征数
        enc_input_dim = self.input_len * self.hist_feat_dim

# Encoder: 把长序列压成一个向量
        self.encoder_in = nn.Linear(enc_input_dim, hidden_size)
        self.encoder_blocks = nn.ModuleList(
            [ResidualBlock(hidden_size, hidden_size * 2) for _ in range(enc_layers)]
        )

# Decoder: 把 latent 映射到 “output_len * decoder_dim”
        decoder_dim = hidden_size
        self.decoder_in = nn.Linear(hidden_size, decoder_dim)
        self.decoder_blocks = nn.ModuleList(
            [ResidualBlock(decoder_dim, decoder_dim * 2) for _ in range(dec_layers)]
        )

# Temporal decoder:逐时间步融合未来协变量
        self.temporal_in = nn.Linear(decoder_dim + fut_cov_dim, temporal_hidden)
        self.temporal_out = nn.Linear(temporal_hidden, 1)

# 线性残差:历史 y -> 未来 y
        self.linear_residual = nn.Linear(self.input_len, self.output_len)

defforward(self, hist_y, hist_cov=None, fut_cov=None):
"""
        hist_y:   (batch, input_len, 1)
        hist_cov: (batch, input_len, hist_cov_dim) or None
        fut_cov:  (batch, output_len, fut_cov_dim) or None
        """

        B = hist_y.size(0)

if hist_cov isnotNone:
            x_hist = torch.cat([hist_y, hist_cov], dim=-1)  # (B, L_in, feat)
else:
            x_hist = hist_y

# flatten 成一个大向量
        x_enc = x_hist.reshape(B, -1)          # (B, L_in * feat_dim)
        x_enc = F.relu(self.encoder_in(x_enc)) # (B, hidden)
for block in self.encoder_blocks:
            x_enc = block(x_enc)

# decoder 得到每个未来时间点的“共享表示”
        x_dec = F.relu(self.decoder_in(x_enc))  # (B, dec_dim)
for block in self.decoder_blocks:
            x_dec = block(x_dec)

# 拓展到每个时间步
        x_dec = x_dec.unsqueeze(1).repeat(1, self.output_len, 1)  # (B, L_out, dec_dim)

# 融合未来协变量
if fut_cov isnotNone:
            td_in = torch.cat([x_dec, fut_cov], dim=-1)
else:
            td_in = x_dec

        h = F.relu(self.temporal_in(td_in))
        y_hat = self.temporal_out(h).squeeze(-1)  # (B, L_out)

# 线性残差(只用历史 y)
        linear_part = self.linear_residual(hist_y.squeeze(-1))  # (B, L_out)

return y_hat + linear_part

训练伪代码也很朴素,大概就是这样:

model = MiniTiDE(
    input_len=96,
    output_len=24,
    hist_cov_dim=4,   # 比如: 小时、星期几、是否周末、是否节假日
    fut_cov_dim=4,
)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.MSELoss()

for batch in dataloader:  # 你自己提前把窗口切好
    hist_y, hist_cov, fut_cov, target = batch
    pred = model(hist_y, hist_cov, fut_cov)
    loss = loss_fn(pred, target)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

这里的 dataloader 就是一个普通的“滑动窗口”数据集: 每个样本拿过去 input_len 个点当输入,后面 output_len 个点当 label,同时准备对应的历史/未来特征即可。真正上生产你可以直接用上面提到的 Darts / neuralforecast / PyTorch Forecasting 这些 TiDE 封装库,它们把窗口、协变量、batching 都帮你搞好了。

4. 什么时候值得上 TiDE?

粗暴给你几个我觉得比较适合的场景:

  • 长时间窗口 + 长预测:比如用过去 30 天预测未来 14 天,这种 Transformer 经常爆显存的场景。
  • 节假日/活动强驱动的业务:有挺多“已知的未来特征”,比如大促、放假、气温预报,这些都可以喂给 temporal decoder。
  • 不想调复杂结构:团队 Python + PyTorch 基础还不错,但是对注意力、RNN 那套没太多经验时,TiDE 这种“全 MLP”的结构其实更好落地。

如果你手头刚好有一条比较干净的业务时间序列,可以试试先用上面这个 MiniTiDE 跑一版 baseline,再对比一下简单线性模型 / LSTM / Transformer 的效果,通常你会发现:TiDE 这种“看起来很普通的 MLP”,在时间序列上还挺能打的。

🔥虎哥私藏精品🔥

虎哥作为一名老码农,整理了全网最全《python高级架构师资料合集》,总量高达650GB