:基于 Koopman 理论的时序分布漂移鲁棒预测实战指南)
人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载KNFKoopman Neural Forecaster是 Google Research 开源仓库中基于 Koopman 理论构建的深度时间序列预测模型专门用于应对现实世界中频繁出现的时间分布漂移temporal distributional shifts。本文以仓库 KNF/README.md 为主线结合 KNF/modules/models.py、KNF/args.py 等源码完整讲解 KNF 的理论动机、模型架构、数据准备、训练脚本、超参数调优以及如何将其迁移到全新数据集上帮助读者从零复现论文实验并落地自己的预测任务。一、背景为什么需要 Koopman 理论来应对分布漂移现实世界的时间序列底层动态往往随时间变化这种时间分布漂移对深度神经网络DNN构成根本性挑战模型在训练分布上拟合得再好一旦动态发生变化预测精度就会急剧下降。KNF 的核心思路是借鉴动力系统领域的 Koopman 理论——在无限维空间中非线性动力系统可以被表示为线性算子Koopman 算子的作用。KNF 通过 DNN 学习一个线性 Koopman 空间以及所选测量函数measurement functions的系数从而将非线性、非平稳的时间序列演化转化为线性动态获得对分布漂移更强的鲁棒性。据仓库 README 所述这是 Koopman 理论首次被应用于没有已知支配方程governing laws的真实混沌时间序列。为了应对动态的变化KNF 引入了三重归纳偏置inductive biases全局算子global operator学习所有序列共享的、稳定的动态特征局部算子local operator通过 Transformer 注意力机制捕获随时间变化的动态反馈回路feedback loop基于回看窗口lookback window上的预测误差持续更新学习到的算子以适配快速变化的行为。论文成果发表于 ICLR 2023见文末引用在多个存在分布漂移的时间序列数据集上展现了优于其他替代方案的预测性能。二、环境依赖与安装仓库通过requirements.txt声明运行环境核心依赖版本如下见 KNF/requirements.txt依赖版本python3.10.8pytorch1.13.1numpy1.23.5absl-py1.2.0cudatoolkit12.0conda23.1.0安装命令pip install -r requirements.txt需要说明的是requirements.txt中同时列出了cudatoolkit与conda等 conda 生态组件pip无法直接安装这两项实际部署建议先通过 Anaconda/Miniconda 创建 Python 3.10 环境并安装 CUDA 12.0 工具链再以pip install -r requirements.txt补齐其余依赖。训练脚本默认优先使用 CUDAdevice torch.device(cuda if torch.cuda.is_available() else cpu)无 GPU 环境会自动回退到 CPU。三、仓库结构总览KNF 目录结构清晰职责分明KNF/ ├── data/ # 数据下载与预处理脚本 │ ├── M4/m4_data_gen.py # 下载 M4 数据并生成训练/测试集 │ ├── Cryptos/cryptos_data_gen.py # 生成加密货币训练/测试集 │ ├── PlayerTraj/traj_data_gen.py # 生成篮球轨迹训练/测试集 │ ├── sample_data/ # M4-weekly 小样本子集含 train.npy/test.npy │ └── data_analysis.py # 可预测性、趋势、季节性分析工具 ├── modules/ # KNF 的 PyTorch 实现 │ ├── data_classes.py # M4 / Cryptos / PlayerTraj / 自定义数据集类 │ ├── train_utils.py # 训练与评估函数 │ ├── normalizer.py # 可逆实例归一化RevIN实现 │ ├── models.py # KNF 核心模型 │ └── eval_metrics.py # 三个数据集对应的三种评估指标 ├── run_koopman.py # KNF 训练主脚本 ├── args.py # 全部超参数定义 ├── run.sh # 在 sample_data 小样本上快速训练 ├── run_exp.sh # 在 M4 / Cryptos / PlayerTraj 全量数据上训练 └── evaluation.py # 对三个数据集进行整体评估四、数据集准备与预处理KNF 支持三类实验数据各自有独立的预处理脚本。4.1 M4 竞赛数据M4 是著名的时序预测竞赛数据集仓库脚本会自动从其官方仓库下载数据并生成训练/测试集python data/M4/m4_data_gen.py生成的数据以 npy 形式保存其中训练文件内部按频率freq组织M4Dataset在加载时通过self.train_data.item().get(freq)取出对应频率如 Weekly、Daily的序列列表因此训练时需通过--data_freq指定频率。4.2 Cryptos 加密货币数据需要先从 Kaggle 的 g-research-crypto-forecasting 竞赛页面下载train.csv和asset_details.csv放到data/Cryptos目录然后运行预处理python data/Cryptos/cryptos_data_gen.py该数据包含 14 种加密货币的多维特征--num_feats8评估时采用加权 RMSEWeighted RMSE权重取自竞赛官方asset_details.csv中各资产的重要性权重权重数组硬编码在 KNF/modules/eval_metrics.py 的WRMSE函数中。4.3 PlayerTraj 篮球运动员轨迹数据需要先下载 NBA 篮球运动员轨迹数据解压所有.7z文件到data/PlayerTraj/json_data然后运行python data/PlayerTraj/traj_data_gen.pyREADME 特别提醒由于采样轨迹时未固定随机种子若要精确复现论文结果请下载论文作者使用的同一份轨迹数据集该数据集与训练好的模型、预测文件一起提供。轨迹数据的特征是二维速度分量--num_feats2评估指标为常规 RMSE。4.4 小样本演示数据仓库自带data/sample_data/train.npy与data/sample_data/test.npy是 M4-weekly 数据的一小部分子集用于快速验证代码流程能否跑通无需任何外部下载。五、KNF 核心模型架构深度解析模型主体实现在 KNF/modules/models.py 的Koopman类中。整体流程可以概括为编码 → 测量 → 线性演化 → 解码。5.1 测量函数Measurement FunctionsKNF 学习显式的测量函数字典dictionary将原始观测映射到 Koopman 空间。默认num_poly/num_sins/num_exp取 -1 时使用多项式函数num_poly3即 1 次、2 次、3 次幂指数函数num_exp1即exp(x)正弦/余弦函数对num_sins input_length // 2 - 1对其余维度为纯数据驱动data-driven的测量函数。对于多变量序列num_feats 1还会计算二阶交互项itertools.combinations生成的所有特征两两组合乘积以捕获特征间的耦合关系。从源码可见KNF/modules/models.py测量函数输入是编码器系数与原观测的逐元素乘积再经多项式、指数、三角函数变换后得到嵌入embedding随后由解码器重建观测。5.2 编码器与解码器编码器MLP输入为每个编码步的观测input_dim * num_feats输出维度为(latent_dim num_sins * 2) * input_dim * num_feats即同时学习测量函数的频率和幅度系数见 KNF/modules/models.py。解码器MLP将 Koopman 空间嵌入映射回观测空间output_dim * num_feats负责重建回看窗口并输出未来预测。两者均为多层 MLP默认 5 层、隐层 256 维支持可选的 InstanceNorm 与 Dropout。5.3 全局算子与局部算子全局 Koopman 算子一个不带 bias 的线性层nn.Linear(latent_dim * num_feats len_interas, ...)对所有时间序列共享学习稳定不变的动态特征默认开启--add_global_operator控制。局部 Koopman 算子由 Transformer 编码器 多头注意力生成。注意力权重被解释为局部线性变换forw torch.einsum(bnl, blh - bnh, forw, local_transform)用于捕获随时间演化的局部动态。5.4 反馈回路Control 模块这是 KNF 应对快速变化行为的关键机制默认开启--add_control控制。源码逻辑为先在整个回看窗口上生成预测inp_preds计算其与真实观测的差异pred_diff该差异输入 Control MLP输出对 Koopman 算子的调整量线性修正矩阵在向前预测时使用修正后的算子local_transform linear_adj进行迭代演化见 KNF/modules/models.py。也就是说模型会根据自身在历史窗口上的表现自适应修正动态算子从而紧跟分布变化。5.5 可逆实例归一化RevINKNF/modules/normalizer.py 提供了可逆实例归一化Reversible Instance Normalization的 PyTorch 实现。模型在编码前对输入做归一化消除序列自身的均值/方差偏移预测后做逆归一化还原到原始尺度。forward方法支持norm与denorm两种模式并可选用仿射参数affineTrue。通过--use_revin默认 True与--use_instancenorm控制是否启用。5.6 自回归多步预测模型一次前向产生num_steps步预测然后通过自回归拼接将预测结果反归一化后拼回输入窗口丢弃最旧的num_steps步重复直到覆盖整个预测目标长度见 KNF/modules/models.py。forward中要求input_length必须能被input_dim整除否则会抛出 ValueError。5.7 秩正则化可选开启--regularize_rank后模型会对局部算子的秩进行正则化通过可学习的dynamics_summary参数构造协方差矩阵计算其特征值平方和与最大值之差作为正则项使最小特征值趋近于 0抑制算子的退化见 KNF/modules/models.py。默认关闭。六、超参数全解析所有超参数集中在 KNF/args.py使用 absl flags 定义训练时通过命令行--flagvalue传入。下表为完整参数清单及默认值参数默认值说明seed123随机种子datasetM4数据集类M4 / Cryptos / Traj / minidata_dirdata_prep/M4/含 train.npy 与 test.npy 的数据目录num_feats1特征数量regularize_rankFalse是否对动态模块做秩正则化use_revinTrue是否使用可逆实例归一化use_instancenormTrue是否对隐状态做实例归一化add_global_operatorTrue是否使用全局 Koopman 算子add_controlTrue是否使用控制反馈模块data_freqNone时间序列频率M4 数据按频率组织learning_rate0.001初始学习率dropout_rate0.0Dropout 比率decay_rate0.9学习率逐 epoch 衰减率StepLRbatch_size128批大小latent_dim64Koopman 潜空间维度num_steps5单次自回归调用预测的步数control_hidden_dim64控制修正矩阵模块隐层维度num_layers5编码器与解码器层数control_num_layers3控制模块层数jumps5生成滑动窗口样本时跳过的步数num_epochs1000最大训练轮数min_epochs60最少训练轮数早停保护input_dim5编码器每一步读取的历史观测数input_length45学习 Koopman 算子的回看窗口长度hidden_dim256编码器/解码器隐层维度train_output_length10训练时反向传播的输出长度test_output_length13测试集预测视界num_heads1Transformer 注意力头数transformer_dim128Transformer 前馈网络维度transformer_num_layers3Transformer 编码器层数num_sins-1正弦/余弦测量函数对数-1 用默认值num_poly-1多项式测量函数最高阶-1 用默认值num_exp-1指数测量函数个数-1 用默认值从 KNF/run_koopman.py 的源码可以看到关键联动关系output_dim被设置为等于input_dim训练时loss_fun nn.MSELoss()优化器为 Adam配 StepLR 学习率调度每 epoch 乘以decay_rate每个样本的损失由四部分构成见 KNF/modules/train_utils.py未来预测损失、回看窗口重建损失、回看窗口单步预测损失、嵌入空间预测损失开启秩正则化时再加上正则项训练时会对梯度做 max_norm5.0 的裁剪防止梯度爆炸验证损失连续 10 个 epoch 均值高于前 10 个 epoch 均值且已训练超过min_epochs时提前停止训练过程中始终保存验证集最优模型.pth若同名模型已存在则自动断点续训resume training。七、训练实战7.1 小样本快速验证run.sh在 M4-weekly 小样本子集上训练一个微型 KNF用于验证代码与流程sh run.sh该脚本内部执行的命令为见 KNF/run.shCUDA_VISIBLE_DEVICES0 python3 run_koopman.py --seed1 --datasetmini \ --data_dirdata/sample_data/ --hidden_dim64 --num_layers3 \ --latent_dim32 --learning_rate0.005 --batch_size16 \ --num_sins10 --transformer_dim64注意这里使用--datasetmini对应 KNF/modules/data_classes.py 中的CustomDataset评估指标为 sMAPE。7.2 全量数据集训练run_exp.sh在三个数据集上并行训练多个 KNF 模型每行一个 GPU 任务末尾wait等待全部完成。各数据集的代表性配置如下M4-Weekly5 个不同种子的模型取预测平均报告 sMAPECUDA_VISIBLE_DEVICES0 python3 run_koopman.py --seed901 --data_freqWeekly \ --datasetM4 --data_dirdata/M4/ --train_output_length10 \ --test_output_length13 --input_dim5 --input_length45 --hidden_dim256 \ --num_layers5 --latent_dim64 --learning_rate0.005 --batch_size128 \ --jumps3 --decay_rate0.85M4-Daily单次运行即可达到当时 SOTA无需集成CUDA_VISIBLE_DEVICES1 python3 run_koopman.py --seed6 --data_freqDaily \ --datasetM4 --data_dirdata/M4/ --train_output_length6 \ --test_output_length14 --input_dim3 --input_length18 --hidden_dim128 \ --num_layers4 --transformer_dim128 --control_hidden_dim64 \ --latent_dim8 --learning_rate0.005 --batch_size256 --jumps5 \ --decay_rate0.85 --num_steps6 --num_sins2PlayerTraj5 次运行报告均值 ± 标准差CUDA_VISIBLE_DEVICES2 python3 run_koopman.py --seed0 --datasetTraj \ --data_dirdata/PlayerTraj/ --num_feats2 --train_output_length15 \ --input_dim3 --input_length21 --hidden_dim128 --latent_dim32 \ --num_layers4 --control_num_layers3 --control_hidden_dim64 \ --transformer_dim64 --transformer_num_layers3 --batch_size128 \ --test_output_length30 --num_steps15 --jumps2 --learning_rate0.001 \ --num_sins3Cryptos5 次运行报告均值 ± 标准差CUDA_VISIBLE_DEVICES0 python3 run_koopman.py --seed162 --datasetCryptos \ --data_dirdata/Cryptos/ --num_feats8 --train_output_length14 \ --test_output_length15 --input_dim7 --input_length63 --hidden_dim64 \ --num_layers5 --latent_dim16 --transformer_num_layers3 \ --transformer_dim256 --control_num_layers3 --control_hidden_dim128 \ --learning_rate0.005 --batch_size512 --jumps100 --num_sins6 --num_steps7从这些配置可以总结出调参规律数据越复杂Cryptos 8 特征、63 回看长度、特征越多latent_dim反而可以更小16而transformer_dim更大256轨迹数据预测视界长30 步num_steps相应增大15以匹配自回归调用。7.3 断点续训与训练产物训练脚本会自动生成dataset_results/目录如M4_results/、Cryptos_results/、Traj_results/模型文件名编码了全部超参数如Koopman_M4_seed901_jumps3_freqWeekly_poly-1_sin-1_exp-1_bz128_lr0.005...pth便于区分不同配置。每次训练结束还会保存test_model_name.pt文件内含test_preds、test_tgts与eval_score三个字段。若同名模型文件已存在脚本会加载并从中断点继续训练见 KNF/run_koopman.py。八、评估与指标8.1 三种评估指标不同数据集使用不同评估指标定义在 KNF/modules/eval_metrics.py映射关系见 KNF/run_koopman.py数据集指标说明M4 / minisMAPE对称平均绝对百分比误差按短/中/长预测视界分别报告CryptosWeighted RMSE按 14 种加密货币的重要性权重加权且只评估最后一个特征15 分钟 ahead 的残差收益PlayerTrajRMSE常规均方根误差预测二维速度分量8.2 整体评估脚本训练全部完成后运行 KNF/evaluation.py 可以对三个数据集做统一评估对 M4-Weekly将多个test_*.pt文件的预测取平均集成计算整体 sMAPE对 PlayerTraj统计 5 次运行的 RMSE 均值与标准差对 Cryptos统计 5 次运行的 Weighted RMSE 均值与标准差。README 指出训练好的模型与其预测文件即论文中报告数字的来源随论文资料一同发布可自行下载对照复现。九、将 KNF 迁移到全新数据集README 给出了在新数据集上训练 KNF 的四步流程结合源码可进一步明确每步的技术要点准备数据文件将新数据集的训练集与测试集分别保存为两个 npy 文件形状必须为(时间序列条数, 序列长度, 特征数)。源码中的CustomDatasetKNF/modules/data_classes.py正是按这一约定加载的。指定数据路径使用CustomDataset加载数据通过--data_dir指向存放train.npy与test.npy的目录脚本内部会拼接出data_dir/train.npy与data_dir/test.npy。--dataset传mini即会走CustomDataset分支。调整特征数根据新数据的特征维度修改--num_feats默认值在 KNF/args.py 中默认是 1。重点调优超参数README 明确建议优先调优input_dim、input_length、num_steps与train_output_length。这四者决定回看窗口的组织方式、自回归步长与训练目标长度直接影响 Koopman 算子的学习质量。此外需要注意input_length必须能被input_dim整除否则模型前向会抛出 ValueError。补充说明数据集类在训练/验证模式下会先对每条序列做全局标准化保存各序列的ts_means与ts_stds用于反归一化并以jumps为步长生成滑动窗口样本按 90%/10% 划分训练集与验证集测试模式下则直接用序列末尾input_length个观测作为输入、test_output_length作为预测视界详见各 Dataset 类的__getitem__。十、总结KNF 把动力系统中的 Koopman 理论与深度序列模型结合通过全局算子、局部算子和反馈回路三重机制为存在时间分布漂移的时间序列提供了鲁棒预测方案。本仓库提供了从数据预处理、模型训练到评估的完整可复现管线data/负责数据准备modules/承载模型与工具run_koopman.py是训练入口args.py集中管理全部超参数run.sh与run_exp.sh分别对应快速验证与全量实验evaluation.py完成多模型集成评估。新数据集接入只需遵循npy 文件 调整特征数 重点调优四个核心超参数的流程即可。引用如果本仓库对您有帮助请引用其对应论文引用格式取自 KNF/README.mdinproceedings{wang2023koopman, title{Koopman Neural Operator Forecaster for Time-series with Temporal Distributional Shifts}, author{Rui Wang and Yihe Dong and Sercan O Arik and Rose Yu}, booktitle{International Conference on Learning Representations}, year{2023} }赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐Nixpkgs 26.05「Yarara」版本发布解读GCC 15、glibc 2.42、默认语言版本升级与重大不兼容变更迁移指南Nixpkgs 26.05「Yarara」版本发布解读GCC 15、glibc 2.42、默认语言版本升级与重大不兼容变更迁移指南 本文基于 Nixpkgs包管理器操作系统WeChatMsg 完整教程免费把微信聊天记录导出成 HTML、Word、CSV还能生成年度聊天报告WeChatMsg 完整教程免费把微信聊天记录导出成 HTML、Word、CSV还能生成年度聊天报告 WeChatMsg 是一个免费开源的微信聊天记录导出工创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考