尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

N-BEATS实战:用PyTorch实现可解释的时间序列预测

N-BEATS实战:用PyTorch实现可解释的时间序列预测 简介这是一套面向时间序列预测场景的N-BEATS深度学习模型Python实现资源覆盖单变量预测从模型搭建、训练到评估与解释的完整流程适合具备一定Python基础、希望掌握可解释神经网络预测方法的数据开发者和研究人员。压缩包共68个文件体积仅179KB以46个Python源码文件为核心涵盖模型定义、训练/预测脚本、数据加载与工具函数另有10个GIN配置文件用于超参数和实验管理5个Jupyter Notebook提供流量、旅游、电力、M3/M4等数据集上的实战演示并附带Dockerfile、依赖清单与说明文档便于快速搭建复现环境。目前已有637人学习下载。借助该实现读者可以剖析N-BEATS中解释性基函数块与残差块的作用机制理解季节性、趋势分解与残差捕捉的实现方式同时还能基于自带数据集或自己的时间序列数据调整参数、完成训练并输出预测结果兼顾了算法原理学习与工程项目落地。1. 没人规定时序模型必须用循环网络N-BEATS这个名字第一次出现时很多人把它读成神经网络版贝茨方法实际上它的全称是 Neural Basis Expansion Analysis for Interpretable Time Series Forecasting。和当时主流的 LSTM、Transformer 类时序模型走的完全不是一条路它不用循环、不用卷积、不用注意力甚至不用位置编码主体结构只有全连接层。但它在 M4 竞赛上把当时的统计方法 SARIMA、ETS 全面压了下去在单变量时序预测上拿到了接近 10% 的精度提升。这个结果在当时让很多习惯用 CNN、RNN 或者图神经网络处理序列的人感到意外也让 N-BEATS 成了深度学习时序预测绕不开的基线模型。本文会从它的核心设计出发用 PyTorch 搭出一个可以训练和推理的精简实现再在公开数据集上把趋势分解、季节分解这类可解释性能力实际跑出来。适合正在做时间序列预测、AI 大作业选题或者想搞清楚为什么堆全连接也能建模时序的 Python 开发者。2. N-BEATS 的核心设计双重残差与基础扩展很多人第一次看 N-BEATS 的结构图会懵一堆方块叠在一起每个方块既输出 backcast 又输出 forecast然后一路减一路加。这个设计和常规的输入整段序列、输出未来一段预测的端到端网络差别很大要先把它拆开看。2.1 从 backcast 和 forecast 两个方向理解每个 BlockN-BEATS 的基本单元是 Block。每个 Block 做的事情是读入一段历史序列输出两个东西一个是 backcast对输入的回溯拟合一个是 forecast对未来的预测。Block 内部是 4 层全连接加 ReLU 激活最后一层同时分出两个头一个头输出 backcast 的系数另一个头输出 forecast 的系数。这两个系数再分别和基础扩展basis做矩阵乘法映射回真实的序列数值。为什么要同时输出 backcast 而不是只输出 forecast这是整个模型最关键的设计。backcast 会从输入序列中被减掉再传给下一个 Block。也就是说后面的 Block 永远只看前面 Block没解释完的残差部分。这样一来每个 Block 不需要一次性把整个序列的特点全部抓完它只需要抓住一种模式就够了——第一个 Block 抓长期趋势第二个抓周期第三个抓剩下的噪声。forecast 这一侧则会把所有 Block 的预测结果累加起来得到最终的预测。一个形象但不完全严谨的理解是backcast 做减法让每个专家各管一段forecast 做加法把每个专家的判断汇总。这种前向残差 后向累加的组合论文里叫双重残差堆叠double residual stacking它让 N-BEATS 在加深网络时不出现梯度消失也不需要 LayerNorm 或者残差连接那样的额外保护。2.1.1 基础扩展用多项式拟合趋势用傅里叶拟合季节基础扩展是 N-BEATS 名字里 Basis Expansion 的来源它决定了 Block 的输出长什么样。在基础版本里模型使用两种 basis趋势 basis 和季节 basis。趋势 basis 由一组多项式基函数组成比如 1、t、t²、t³……模型通过线性组合这些基函数就能拟合出上升、下降、平滑变化这类整体走势。季节 basis 则由傅里叶基函数组成即一组不同频率的正弦和余弦波通过线性组合它们可以拟合出周期性的模式。这种设计有一个直接的工程好处如果你把预测结果按 basis 拆开能明确区分出哪一部分是趋势贡献的、哪一部分是季节贡献的。这在金融销量预测、容量规划这类需要向业务方解释预测依据的场景里非常实用。相比 LSTM 和 Transformer 的黑盒输出N-BEATS 的输出天然自带分解视图。2.2 Stack 的组织方式Generic 与 Interpretable 变体N-BEATS 论文里给出了两种配置Generic 和 Interpretable。Generic 版本的每个 Block 用随机初始化的 basis 矩阵模型的表达能力更强适合纯精度优先的任务。Interpretable 版本按 Stack 划分职责前几个 Stack 固定用趋势多项式 basis后面几个 Stack 固定用季节傅里叶 basis每个 Stack 内部共享同一组 basis。训练时可以根据数据特点选择配置。如果序列有明显的趋势加季节成分而且你需要向别人说明预测依据用 Interpretable 配置更合适。如果序列模式复杂、没有先验知识可用Generic 配置通常效果更好。一个 Block 的 hidden 层数论文默认 4 层 512 维和 basis 维度趋势 4 阶、季节 8 对傅里叶项这类设置是决定模型容量和可解释粒度的主要旋钮。2.3 为什么它不需要注意力机制N-BEATS 能抛开注意力机制还有一个原因它的感受野天然覆盖整个输入窗口。Transformer 处理长序列时需要注意力矩阵建模任意两个位置的关系N-BEATS 直接把整个窗口压成全连接层的输入理论上每个输出位置都隐式依赖全部输入位置。全连接网络对输入做的是全局混合只不过这种混合是静态的、训练后固定的而注意力是动态的、随输入变化的。对于单变量时序预测静态混合通常已经足够因为模式相对稳定。当然这也暴露了它的一个边界N-BEATS 不容易吸收外生变量特征后期 N-BEATSx 的提出正是为了解决这一问题。3. 用 PyTorch 写一个可运行的 N-BEATS 精简实现理论部分说得再多不如直接落成代码。这一节我会给出一个结构完整、可直接训练的 N-BEATS 实现。它不追求复刻论文的每一个细节但保留了双重残差堆叠、趋势基础扩展、季节基础扩展这三个核心机制。使用 PyTorch 2.x 和 Python 3.10 均可直接运行不需要额外的特殊依赖。3.1 定义 Block 和基础扩展层先定义最底层的 Block。代码的关键在于把 backcast 和 forecast 系数的生成与 basis 矩阵的映射分离这样后续换 basis 时不用改动 Block 结构import torch import torch.nn as nn import numpy as np class NBEATSBlock(nn.Module): def __init__(self, input_size, hidden_size256, num_layers4, backcast_basisNone, forecast_basisNone): super().__init__() # 4 层全连接最后一层同时分出 backcast 和 forecast 两个头 layers [] in_dim input_size for _ in range(num_layers): layers.append(nn.Linear(in_dim, hidden_size)) layers.append(nn.ReLU()) in_dim hidden_size self.fc nn.Sequential(*layers) self.backcast_theta nn.Linear(hidden_size, input_size) self.forecast_theta nn.Linear(hidden_size, input_size) self.backcast_basis backcast_basis self.forecast_basis forecast_basis def forward(self, x): h self.fc(x) theta_b self.backcast_theta(h) theta_f self.forecast_theta(h) backcast torch.matmul(theta_b, self.backcast_basis.T) forecast torch.matmul(theta_f, self.forecast_basis.T) return backcast, forecast这里有两个容易踩的细节。第一theta_b的维度是(batch, input_size)而后面的 basis 矩阵维度是(input_size, input_size)通过矩阵乘法把系数映射回序列空间。第二backcast_basis和forecast_basis通常使用相同的 basis 矩阵但论文在实现中允许不同例如可以给 forecast 方向添加额外的高频基函数。这里统一使用同一份即可。3.2 趋势与季节基础扩展生成接下来是基础扩展部分。趋势用多项式基函数季节用傅里叶级数。需要留意的是多项式基函数的数值范围随阶数增长会变得非常大因此生成后要做归一化处理否则训练前期容易梯度爆炸。def trend_basis(degree4, size128): t np.arange(size) / size basis np.vstack([t ** i for i in range(degree 1)]) # 逐行做 min-max 归一化保证各阶多项式数值范围一致 basis (basis - basis.min(axis1, keepdimsTrue)) / \ (basis.max(axis1, keepdimsTrue) - basis.min(axis1, keepdimsTrue) 1e-8) return torch.tensor(basis, dtypetorch.float32) def seasonality_basis(num_pairs6, size128): t np.arange(size) / size basis [] for i in range(1, num_pairs 1): basis.append(np.sin(2 * np.pi * i * t)) basis.append(np.cos(2 * np.pi * i * t)) return torch.tensor(np.vstack(basis), dtypetorch.float32)注意trend_basis里size是输入窗口的长度必须和输入序列被展平后的长度一致。seasonality 的num_pairs决定能建模到多高频的周期成分1 对应一个完整周期2 对应半周期依此类推。对于以日为单位的数据如果周期是 24 小时而窗口长度是 96 点把num_pairs设到 4 以上才能覆盖到日周期的高次谐波。另一个值得说明的点是trend_basis生成的是输入窗口内的基函数每个 Block 用它生成的是窗口内各时间点的加权系数。3.3 组装 Stack 与完整模型把多个 Block 通过 backcast 残差连接串联起来就组成了 Stack多个 Stack 之间也是同样的连接方式。论文里的 Stack 内部各 Block 共享 basis我这里为了让代码结构更清晰允许每个 Block 持有自己的 basisGeneric 风格。如果想实现 Interpretable 风格只需要让同一 Stack 内的 Block 使用同一份 basis 矩阵class NBEATS(nn.Module): def __init__(self, input_size128, forecast_size96, stack_num4, block_num_per_stack3, hidden_size256, trend_degree4, season_pairs6): super().__init__() self.forecast_size forecast_size trend_b trend_basis(trend_degree, input_size) seas_b seasonality_basis(season_pairs, input_size) trend_f trend_basis(trend_degree, forecast_size) seas_f seasonality_basis(season_pairs, forecast_size) self.blocks nn.ModuleList() for stack_idx in range(stack_num): for _ in range(block_num_per_stack): if stack_idx stack_num // 2: bb, bf trend_b, trend_f else: bb, bf seas_b, seas_f self.blocks.append( NBEATSBlock(input_size, hidden_size, backcast_basisbb, forecast_basisbf) ) def forward(self, x): # 输入 x 形状: (batch, input_size) forecast_sum torch.zeros(x.shape[0], self.forecast_size, devicex.device) residual x for block in self.blocks: backcast, forecast block(residual) residual residual - backcast forecast_sum forecast_sum forecast return forecast_sumstack_num和block_num_per_stack控制模型的深度。当 stack 数为 4 时前两个负责趋势解释后两个负责季节解释这是 Interpretable 的典型分配。hidden_size是每个 Block 里全连接层的宽度论文默认是 512但实践中 256 在大部分数据集上已经够用还能显著减少显存占用和训练时间。3.3.1 参数规模与初值设置N-BEATS 的参数量主要由input_size * hidden_size * num_layers贡献。例如输入窗口 128、隐藏层 256、每个 Block 4 层全连接则单个 Block 的参数量约 128×256 256×256×3 256×128 ≈ 20 万。整个模型 12 个 Block 约 240 万参数属于中小规模模型一张普通显卡就能跑。需要特别注意的是全连接层的权重初始化默认是均匀分布或正态分布如果输入序列的数值范围很大比如销量数据从 0 到 10 万波动建议在数据预处理中做标准化否则前向传播时theta_b * basis的乘积很容易溢出。3.4 训练循环MAE 损失是时序预测的默认选择N-BEATS 的论文用的是 MAEL1Loss而不是 MSE。原因在于 MAE 对离群点更鲁棒预测结果会更贴近中位数而非均值这对真实业务数据通常更友好。训练代码可以这样写criterion nn.L1Loss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) for epoch in range(100): model.train() total_loss 0 for x, y in train_loader: optimizer.zero_grad() pred model(x) loss criterion(pred, y) loss.backward() # 梯度裁剪防止个别样本导致训练崩溃 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total_loss loss.item() print(fEpoch {epoch}: loss {total_loss / len(train_loader):.6f})lr1e-4配合 Adam 是一个保守但稳定的起点。梯度裁剪的阈值 1.0 在 N-BEATS 这类全连接网络上不是必需的但加上后能避免偶发的离群样本把某一层权重冲偏。如果发现 loss 早期下降过慢可以把 lr 调到 3e-4但要注意观察后续是否出现震荡。4. 用 ETTh1 数据集训练参数设置与踩坑有了模型结构接下来要把它放到真实数据上验证。ETTh1Electricity Transformer Temperature是时序预测领域常用的公开数据集记录了一台电力变压器的油温、负载等指标小时级别采样且包含明显的周期性和趋势。在 N-BEATS-master.zip 这类项目包中你通常能找到类似的示例数据脚本如果没有也可以自己从公开渠道下载 CSV 后按本节流程处理。注意这里不使用任何外生变量纯靠单变量历史序列做预测。4.1 数据准备归一化、滑窗、划分在把数据送入模型之前要做三步处理一是把序列按比例划分训练/验证/测试二是做归一化三是把长序列切成滑窗样本。下面这份代码以 PyTorch 的 Dataset 方式实现能直接配合 DataLoader 使用import pandas as pd import numpy as np from torch.utils.data import Dataset, DataLoader class TimeSeriesDataset(Dataset): def __init__(self, data, input_size128, forecast_size96): self.data data self.input_size input_size self.forecast_size forecast_size def __len__(self): return len(self.data) - self.input_size - self.forecast_size 1 def __getitem__(self, idx): x self.data[idx: idx self.input_size] y self.data[idx self.input_size: idx self.input_size self.forecast_size] return torch.tensor(x, dtypetorch.float32), \ torch.tensor(y, dtypetorch.float32) # 读取 OT 列去除缺失值 df pd.read_csv(ETTh1.csv)[OT].dropna().values # 用训练集的均值/方差做归一化 train_data df[: int(len(df) * 0.7)] mean, std train_data.mean(), train_data.std() data (df - mean) / std dataset TimeSeriesDataset(data, input_size128, forecast_size96) train_len int(len(dataset) * 0.8) train_ds, val_ds torch.utils.data.random_split(dataset, [train_len, len(dataset) - train_len]) train_loader DataLoader(train_ds, batch_size512, shuffleTrue, num_workers4)这里input_size128对应过去 128 小时forecast_size96对应未来 4 天。一个需要说明的工程点是验证集或测试集不参与均值/方差计算否则会造成信息泄漏。常见做法是先只在训练集上算mean和std再应用到整个数据集。4.2 关键参数选择窗口长度、Stack 数、basis 阶数参数设置决定了模型上限。直接给出我常用的参数组合参数取值说明input_size128预测长度的 1.52 倍较好太短会丢失周期信息forecast_size96常用 24/48/96预测越长难度越大stack_num42 个趋势 stack 2 个季节 stackblock_num_per_stack3每个 stack 3 个 block总共 12 层hidden_size256增大到 512 精度提升有限训练时间约翻倍trend_degree4多项式最高阶次过大容易过拟合season_pairs6对应 6 组正余弦能覆盖 16 次谐波batch_size512显存不高也能跑过小会导致梯度噪声大learning_rate1e-4Adam 默认足够稳定lossMAE对离群点鲁棒预测中位数而非均值其中input_size和forecast_size的比例是我反复实验后的经验值。预测长度短时输入长度可以是预测长度的 3 倍以上预测长度长时超过 2 倍容易在窗口尾部引入过多的噪声信息。季节序列的周期必须能被窗口长度整除否则傅里叶基函数在窗口内无法形成整数个周期。4.3 训练观察loss 曲线和过拟合判断训练过程中我会同时在验证集上计算 loss每 10 个 epoch 打印一次。正常训练时训练 loss 和验证 loss 会同步下降最后趋于平稳。如果验证 loss 开始上升而训练 loss 还在下降说明模型开始记住训练集的噪声需要提前停止或增大数据规模。# 后台训练并把日志写入文件方便随时查看 nohup python train.py --input_size 128 --forecast_size 96 \ --stack_num 4 --hidden_size 256 --epochs 100 train.log 21 # 实时查看 loss 变化 tail -f train.log在 ETTh1 上通常 3050 个 epoch 后验证 loss 就不再明显下降。如果你的 loss 在 10 个 epoch 内就出现大幅波动优先检查三件事学习率是否过大、batch 是否太小、数据归一化是否正确。这三个问题占了时序模型训练失败原因的八成以上。4.4 两个容易被忽视的坑第一个坑是忘记shuffleTrue。时序数据按时间顺序排列如果不打乱同一个 batch 内的样本时间高度重叠模型会很快过拟合到一个局部模式且泛化很差。第二个坑是训练/验证划分的随机切分方式。random_split会把时序打乱这在纯预测任务里是可接受的因为在验证集上观测到的分布和训练集一致。但如果你希望模拟用过去预测未来的真实场景需要按时间顺序划分——比如前 80% 的时间戳作训练后 20% 作验证。两种划分方式各有道理前者测分布拟合能力后者测时序外推能力。5. 把预测拆开看N-BEATS 的可解释性验证最后一步把训练好的模型输出按趋势和季节分解。这也是 N-BEATS 与 RNN、Transformer 类模型拉开差距的关键特性。可解释性不仅能让你向业务方交代预测为什么长这样还能帮你诊断模型问题。5.1 用 hook 提取每个 block 的 forecast 分量由于模型中每个 Block 的forecast直接保存了该 Block 对所有未来时间点的贡献最简单的办法是修改模型并记录这些值def get_decomposition(model, input_x): model.eval() residual input_x trend_sum 0 season_sum 0 with torch.no_grad(): for i, block in enumerate(model.blocks): backcast, forecast block(residual) residual residual - backcast # 前一半 stack 的贡献视为趋势 if i len(model.blocks) // 2: trend_sum trend_sum forecast else: season_sum season_sum forecast return trend_sum, season_sum这里把前一半 block 的输出累加为趋势、后一半为季节前提是模型使用了 Interpreable 配置。如果你的模型是 Generic 配置各 block 没有明确分工这样拆分就不具备语义含义。5.2 可视化趋势和季节分量拿到trend_sum和season_sum后可视化时需要注意它们都是形状为(forecast_size,)的一维序列。绘制时把历史窗口的最后 24 个点、趋势分量、季节分量、以及真实未来值画在同一张图上。观察以下三点趋势分量是否平滑无毛刺、季节分量是否呈现规律的峰谷周期、两者相加是否接近最终预测。如果趋势分量出现突兀拐折通常说明趋势多项式的阶数太高可以调低两个度如果季节分量振幅明显小于真实波动说明季节 stack 学习不足可以增加block_num_per_stack。别忘了处理归一化。模型预测的是标准化后的数据可视化时要把mean和std加回去trend trend_sum.numpy() * std mean season season_sum.numpy() * std mean5.3 验证预测一致性残差分析与滚动预测可解释性之后还要做一层数值验证检查逐点残差y_true - y_pred的分布。如果残差均值偏离 0 明显说明模型存在系统性偏置。常见做法是计算残差的自相关系数——如果残差在滞后 24 点处仍有显著相关说明季节模式没有被完全捕捉。此时可以将season_pairs从 6 提升到 8 或更高或者增大输入窗口让模型看到更多完整周期。对于需要周期性滚动预测的生产场景我通常还会做一个简单的滚动回测每预测完一段把真实值拼接到窗口末尾并移除同等长度的最老数据不断向前滚动。N-BEATS 的优势在于推理时无需重新训练只要窗口长度不变滚动预测的开销基本可控。唯一的限制是输入窗口被全连接层固定为训练长度重新改变输入尺寸就需要重新训练模型。这也是它的一个边界所在。本文还有配套的精品资源点击获取
返回列表