NeurIPS 2024 | 用于时间序列预测的检索增强扩散模型

文章信息

文章题目为《Retrieval-Augmented Diffusion Models for Time Series Forecasting》,发表在第38届Conference on Neural Information Processing Systems,作者Jingwei Liu,来自北京大学

摘要

虽然时间序列扩散模型近年来受到了广泛关注,但现有模型的性能仍然高度不稳定。限制时间序列扩散模型的因素主要包括时间序列数据集不足以及缺乏指导机制。为解决这些问题,本文提出了一种检索增强型时间序列扩散模型(Retrieval-Augmented Time series Diffusion,RATD)。RATD的框架由两个部分组成:基于嵌入的检索过程和参考引导的扩散模型。在第一部分中,RATD从数据库中检索出与历史时间序列最相关的时间序列作为参考。在第二部分中,这些参考被用于引导去噪过程。本文的方法能够利用数据库中的有意义样本辅助采样,从而最大化数据集的利用率。同时,这种参考引导机制也弥补了现有时间序列扩散模型在指导方面的不足。在多个数据集上的实验与可视化结果表明,本文的方法在复杂预测任务中尤其有效。

引言

时间序列预测在多种应用中发挥着关键作用,包括天气预测、金融预测、地震预测以及能源规划。解决时间序列预测任务的一种方式是将其视为条件生成任务,即利用条件生成模型来学习条件分布P(xp∣xH),在已知历史序列xH的情况下预测目标时间序列xp。作为当前最先进的条件生成模型,扩散模型已经被广泛应用于时间序列预测任务中。
虽然现有的时间序列扩散模型在部分预测任务上表现良好,但在某些场景下仍然存在不稳定。限制时间序列扩散模型性能的因素较为复杂,其中两个尤其明显。其一,大多数时间序列缺乏直接的语义或标签对应关系,这往往导致时间序列扩散模型在生成过程中缺乏有意义的指导(不同于图像扩散模型中的文本指导或标签指导)。这一点也限制了时间序列扩散模型的潜力。其二,时间序列数据集本身存在规模不足和分布不均衡的问题。与图像数据集相比,时间序列数据集的规模通常要小得多。常见的图像数据集(如LAION-400M)包含4亿个样本对,而大多数时间序列数据集通常只有数万个数据点。在数据规模不足的情况下训练扩散模型以学习精确分布具有挑战性。此外,真实的时间序列数据集还表现出显著的不平衡。例如,在现有的心电图数据集MIMIC-IV中,与诊断为预激综合征(PS)相关的记录占总记录数的比例不足0.025%。这种不平衡现象可能导致模型忽略极其稀有的复杂样本,训练过程中更倾向于生成更常见的预测,从而难以处理复杂预测任务。
为了解决这些限制,本文提出了检索增强型时间序列扩散模型(Retrieval-Augmented Time series Diffusion,RATD),专门用于复杂的时间序列预测任务。本文的方法由两部分组成:基于嵌入的检索和参考引导的扩散模型。在获得历史时间序列后,首先将其输入基于嵌入的检索过程,检索出与其最相似的k个样本作为参考。这些参考样本会在去噪过程中作为指导。RATD的核心在于通过寻找数据集中与历史时间序列最相关的参考样本来最大化现有数据集的利用率,从而为去噪过程提供有意义的指导。RATD不仅提高了有限时间序列数据的利用效率,而且在一定程度上缓解了数据不平衡所带来的问题。同时,这种参考引导机制也
弥补了现有时间序列扩散模型在指导方面的不足。本文的方法在多个数据集上表现出强劲的性能,尤其是在更复杂的任务中。
总结而言,本文的主要贡献如下:
(1)为解决复杂的时间序列预测问题,本文首次提出了检索增强型时间序列扩散模型(RATD),该方法能够更充分地利用数据集,并在去噪过程中提供有意义的指导。
(2)本文设计了额外的参考调制注意力(Reference Modulated Attention,RMA)模块,在去噪过程中从参考样本中提供合理的指导。RMA能够有效且简洁地整合信息,同时不会引入过多额外的计算开销。
(3)本文在五个真实世界的数据集上进行了实验,并基于多种指标提供了全面的结果展示与分析。实验结果表明,本文的方法在多个任务中取得了与基线方法相当或更优的表现。

模型构建

下图展示了RATD的整体架构。本文基于DiffWav搭建了整个流程,该模型结合了传统的扩散模型框架和二维Transformer结构。在预测任务中,RATD首先会根据输入的历史事件序列,从数据库DR中检索相关序列。这些检索到的样本随后会作为参考输入到参考调制注意力(Reference-Modulated Attention, RMA)模块。在RMA层中,本文将时间步t的输入特征[xH,xt]与辅助信息Is及参考序列xR进行融合。通过这种融合,参考序列能够引导生成过程。本文将在接下来的小节中介绍这些流程。

(1).时间序列检索数据库的构建
在进行检索之前,需要先构建合适的数据库。本文提出了一种针对不同特征的时间序列数据集的数据库构建策略。某些时间序列数据集规模不足,且难以用单一类别标签进行标注(如电力时间序列);而另一些数据集则包含完整的类别标签,但存在显著的类别不平衡问题(如医学时间序列)。针对这两类数据集,本文使用了两种不同的数据库定义方式:
方式一:将整个训练集直接定义为数据库DR::

其中xi={si,⋯,si+l+h}为长度为l+h的时间序列,Dtrain表示训练集。
方式二:将包含所有类别样本的子集定义为数据库D’R:

其中xik是训练集中第k类的第i个样本,长度为l+h,C是原始数据集的类别集合。
为简便起见,本文将两类数据库都记作DR。
(2).检索增强型时间序列扩散
基于嵌入的检索机制
在时间预测任务中,理想的参考序列{si,⋯,si+h}应该是前n个点{si-n,⋯,si-1}与数据库DR中历史时间序列{sj,⋯,sj+h}最相关的样本。在本文的方法中,更关注时间序列整体的相似性。本文通过嵌入之间的距离来量化时间序列间的参考关系。
为了确保嵌入能够有效表示整个时间序列,本文使用了预训练编码器Eφ在表示学习任务上进行训练,在本文的检索机制中其参数φ是冻结的。对于数据库DR中长度为n+h的时间序列,仅前n个点被编码,因此数据库可表示为:

其中[p:q]表示时间序列中从第p到第q的子序列。历史时间序列对应的嵌入表示为vH= Eφ(xH)。随后本文计算vH与所有嵌入的距离,并检索出距离最小的k个参考序列:

由此,本文基于查询xH得到数据库DR的一个子集xR,即,其中∣xR∣=k。
参考引导的时间序列扩散模型
在扩散过程中,前向过程与传统扩散过程相同。反向过程的目标是推断后验分布p(ztar∣zc):

其中p(xT∣xH)≈N(xT∣xH,I),pθ(xt−1∣xt,xH,xR)是从xt到xt−1的可学习反向转移核。通常假设:

其中μθ是带参数θ的深度神经网络。反向过程中的Σθ被近似为固定值。因此,通过设计合理且鲁棒的μθ,即可实现参考引导的去噪。
(3).去噪网络架构
与DiffWave和CSDI类似,本文的网络基于Transformer层构建。但现有框架无法有效利用参考作为指导。考虑到通过注意力机制融合xR与xt的直觉,本文提出了新模块参考调制注意力(RMA)。
与普通注意力模块不同,RMA融合了三类特征:当前时间序列特征、辅助特征和参考特征。具体来说,RMA被设置在每个残差模块的开头,具体网络结构见下图。本文使用1D-CNN 从输入xt、参考xR和辅助信息中提取特征(所有参考会被拼接后再提取特征)。辅助信息由两部分组成,表示当前时间序列数据集中变量与时间步之间的相关性。本文通过线性层调整三类特征的维度,并通过矩阵点积进行融合。类似于文本-图像扩散模型,RMA 可以有效利用参考信息引导去噪,同时合适的参数设置能够防止结果过度依赖参考。

为了训练 RATD(即优化其引入的证据下界),本文采用与先前工作相同的目标函数。时间步t−1的损失定义如下:

其x^0是由xt预测得到的

是扩散过程中的超参数。

实验设置

数据集
本文在四个常用的真实世界时间序列数据集上进行了实验:
Electricity:包含321位客户两年内的逐小时电力消耗数据;
Wind:包含2020–2021年的风力发电记录;
Exchange:记录了八个国家(澳大利亚、英国、加拿大、瑞士、中国、日本、新西兰和新加坡)的每日汇率;
Weather†:记录了2020–2021年期间,以10分钟间隔采集的21项气象指标。
此外,本文还将所提方法应用于一个大型心电图(ECG)时间序列数据集:MIMIC-IV-ECG。该数据集包含在Beth Israel Deaconess 医疗中心(BIDMC)采集的超过19万名患者和45万次住院的临床心电图数据。
基线方法
为了全面展示本文方法的有效性,本文将RATD与四类时间序列预测方法进行对比。基线方法包括:
时间序列扩散模型:CSDI、mr-Diff、D3VAE、TimeDiff;
结合频率信息的最新时间序列预测方法:FiLM、Fedformer、FreTS;
时间序列Transformer方法:PatchTST、Autoformer、Pyraformer、Informer、iTransformer;
其他常见方法:TimesNet、SciNet、Nlinear、DLinear、NBeats。
评估指标
为了全面评估本文提出的方法,本实验采用三类指标:
概率预测指标:在每个时间序列维度上计算连续分级概率得分(Continuous Ranked Probability Score, CRPS);
距离指标:采用均方误差(Mean Squared Error, MSE)和平均绝对误差(Mean Absolute Error, MAE)来衡量预测结果与真实值之间的差异。

实验结果

下表展示了本文在四个日级数据集上的主要实验结果。本文的方法在性能上超越了现有的时间序列扩散模型。与其他时间序列预测方法相比,本文的方法在四个数据集中的三个上表现优异,在剩余的一个数据集上也具有竞争力。值得注意的是,本文在Wind数据集上取得了突出成果。由于该数据集缺乏明显的短期周期性(如日周期或小时周期),其中部分预测任务对其他模型而言极具挑战,而检索增强机制能够有效帮助解决这些困难的预测任务。

下图展示了本文在Wind数据集实验中随机选取的一个案例研究。本文将预测结果与iTransformer以及两个流行的开源时间序列扩散模型CSDI和D3VAE进行了对比。尽管CSDI和D3VAE在短期预测初期表现准确,但由于缺乏指导,其长期预测与真实值出现了明显偏差。iTransformer能够捕捉大致趋势和周期模式,但本文的模型在预测质量上优于所有对比方法。此外,从图中预测结果与参考序列的对比可以看出,虽然参考序列提供了强有力的指导,但并未直接替代完整的生成结果,这进一步验证了本文方法的合理性。

下表展示了本文的方法在MIMIC-IV-ECG数据集上的测试结果。本文选择了一些性能较强的开源方法作为基线进行比较。实验分为两部分:第一部分评估整个测试集;第二部分则从测试集中选取稀有病例(占比不足2%的样本)作为子集进行评估。第二部分的预测任务对深度模型来说更具挑战性。在第一项实验中,本文的方法取得了接近iTransformer的结果;而在第二项任务中,本文的模型显著优于其他方法,充分体现了其在处理复杂挑战性任务时的有效性。


结论

本文提出了一种新的时间序列扩散建模框架,以解决现有扩散模型在预测性能上的局限性。RATD 从构建的数据库中检索与历史时间序列最相关的样本,并将其作为参考来引导扩散模型的去噪过程,从而获得更加准确的预测结果。在五个真实世界数据集上的实验结果表明,RATD 在应对复杂的时间序列预测任务时表现出了极高的有效性。

欢迎关注微信公众号《当交通遇上机器学习》!如果你和我一样是轨道交通、道路交通、城市规划相关领域的,也可以加微信:Dr_JinleiZhang,备注“进群”,加入交通大数据交流群!希望我们共同进步!
往
期
回
顾
重磅发布 | 《Artificial Intelligence for Transportation》新刊上线!
团队研究成果|大型活动期间城轨短时客流预测
团队研究成果|城市轨道交通短时交通客流OD预测
团队研究成果|城市轨道交通新线客流预测
团队研究成果|城市轨道交通疫情期间短时客流预测
团队研究成果|基于物理信息引导的突发事件期间的城市轨道交通短时OD需求预测
我知道你在看哟

更多推荐
所有评论(0)