码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • PyTorch搭建Transformer实现多变量多步长时间序列预测(负荷预测)


    目录

    • I. 前言
    • II. Transformer
    • III. 代码实现
      • 3.1 数据处理
      • 3.2 模型训练/测试
      • 3.3 实验结果
    • IV. 源码及数据

    I. 前言

    前面已经写了很多关于时间序列预测的文章:

    1. 深入理解PyTorch中LSTM的输入和输出(从input输入到Linear输出)
    2. PyTorch搭建LSTM实现时间序列预测(负荷预测)
    3. PyTorch搭建LSTM实现多变量时间序列预测(负荷预测)
    4. PyTorch搭建双向LSTM实现时间序列预测(负荷预测)
    5. PyTorch搭建LSTM实现多变量多步长时间序列预测(一):直接多输出
    6. PyTorch搭建LSTM实现多变量多步长时间序列预测(二):单步滚动预测
    7. PyTorch搭建LSTM实现多变量多步长时间序列预测(三):多模型单步预测
    8. PyTorch搭建LSTM实现多变量多步长时间序列预测(四):多模型滚动预测
    9. PyTorch搭建LSTM实现多变量多步长时间序列预测(五):seq2seq
    10. PyTorch中实现LSTM多步长时间序列预测的几种方法总结(负荷预测)
    11. PyTorch-LSTM时间序列预测中如何预测真正的未来值
    12. PyTorch搭建LSTM实现多变量输入多变量输出时间序列预测(多任务学习)
    13. PyTorch搭建ANN实现时间序列预测(风速预测)
    14. PyTorch搭建CNN实现时间序列预测(风速预测)
    15. PyTorch搭建CNN-LSTM混合模型实现多变量多步长时间序列预测(负荷预测)
    16. PyTorch搭建Transformer实现多变量多步长时间序列预测(负荷预测)

    上述文章中都没有涉及到近些年来比较火的Attention机制,随Attention机制一起提出的是transformer模型,关于transformer模型的原理网上各种讲解很多,这里就不具体描述了,有机会再写。

    II. Transformer

    PyTorch封装了Transformer的具体实现,如果导入失败可以参考:torch.nn.Transformer导入失败。

    Transformer模型搭建如下:

    class TransformerModel(nn.Module):
        def __init__(self, args):
            super(TransformerModel, self).__init__()
            self.args = args
            # embed_dim = head_dim * num_heads?
            self.input_fc = nn.Linear(args.input_size, args.d_model)
            self.output_fc = nn.Linear(args.input_size, args.d_model)
            self.pos_emb = PositionalEncoding(args.d_model)
            encoder_layer = nn.TransformerEncoderLayer(
                d_model=args.d_model,
                nhead=8,
                dim_feedforward=4 * args.input_size,
                batch_first=True,
                dropout=0.1,
                device=device
            )
            decoder_layer = nn.TransformerDecoderLayer(
                d_model=args.d_model,
                nhead=8,
                dropout=0.1,
                dim_feedforward=4 * args.input_size,
                batch_first=True,
                device=device
            )
            self.encoder = torch.nn.TransformerEncoder(encoder_layer, num_layers=8)
            self.decoder = torch.nn.TransformerDecoder(decoder_layer, num_layers=8)
            self.fc = nn.Linear(args.output_size * args.d_model, args.output_size)
    
        def forward(self, x, y):
            # print(x.size())  # (256, 24, 7)
            x = self.input_fc(x)  # (256, 24, 128)
            x = self.pos_emb(x)   # (256, 24, 128)
            x = self.encoder(x)
            # print(y.size())   # (256, 4, 7)
            y = self.output_fc(y)   # (256, 4, 128)
            out = self.decoder(y, x)  # (256, 4, 128)
            out = out.view(out.shape[0], -1)   # (256, 4 * 128)
            out = self.fc(out)  # (256, 4)
    
            return out
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20
    • 21
    • 22
    • 23
    • 24
    • 25
    • 26
    • 27
    • 28
    • 29
    • 30
    • 31
    • 32
    • 33
    • 34
    • 35
    • 36
    • 37
    • 38
    • 39
    • 40

    初始时的数据输入维度为7,也就是每个时刻的负荷值以及6个环境变量。在Transformer的原始论文中,文本的嵌入维度为512,而且PyTorch规定nhead数和d_model也就是嵌入维度必须满足整除关系,因此首先将原始数据从7维映射到d_model维度:

    x = self.input_fc(x)
    
    • 1

    其中input_fc:

    self.input_fc = nn.Linear(args.input_size, args.d_model)
    
    • 1

    然后对原始输入进行位置编码:

    x = self.pos_emb(x)
    
    • 1

    然后经过编码层:

    x = self.encoder(x)
    
    • 1

    得到的输出和输入维度一致。

    接着将编码器输出x和标签y同时输入解码器进行解码:

    y = self.output_fc(y)   # (256, 4, 128)
    out = self.decoder(y, x)
    
    • 1
    • 2

    标签y在进入解码器前同样需要将其维度由7映射到d_model。

    值得注意的是,在前面的文章中,y的维度都是(batch_size, output_size),而在Transformer中,y的维度为(batch_size, output_size, d_model)。

    III. 代码实现

    3.1 数据处理

    利用前24小时的负荷值+环境变量预测后4个时刻的负荷值,数据处理和前面一致,只是需要注意的是,y中不再只含有负荷值这1个变量,而是和x一样,都含有7个变量。

    3.2 模型训练/测试

    和前文一致。

    3.3 实验结果

    相关参数如下所示:

    def args_parser():
        parser = argparse.ArgumentParser()
    
        parser.add_argument('--epochs', type=int, default=50, help='input dimension')
        parser.add_argument('--seq_len', type=int, default=24, help='seq len')
        parser.add_argument('--input_size', type=int, default=7, help='input dimension')
        parser.add_argument('--d_model', type=int, default=128, help='input dimension')
        parser.add_argument('--output_size', type=int, default=4, help='output dimension')
        parser.add_argument('--lr', type=float, default=2e-4, help='learning rate')
        parser.add_argument('--batch_size', type=int, default=256, help='batch size')
        parser.add_argument('--optimizer', type=str, default='adam', help='type of optimizer')
        parser.add_argument('--device', default=torch.device("cuda" if torch.cuda.is_available() else "cpu"))
        parser.add_argument('--weight_decay', type=float, default=1e-9, help='weight decay')
        parser.add_argument('--bidirectional', type=bool, default=False, help='LSTM direction')
        parser.add_argument('--step_size', type=int, default=10, help='step size')
        parser.add_argument('--gamma', type=float, default=0.5, help='gamma')
    
        args = parser.parse_args()
    
        return args
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9
    • 10
    • 11
    • 12
    • 13
    • 14
    • 15
    • 16
    • 17
    • 18
    • 19
    • 20

    训练50轮,MAPE为5.04%:
    在这里插入图片描述

    IV. 源码及数据

    基于PyTorch的Transformer时间序列预测代码

  • 相关阅读:
    PHP(3)PHP基础语法
    图片如何转换成PDF格式?教你一招快速转换
    安卓使用okhttpfinal下载文件,附带线程池下载使用
    Sentinel源码剖析之核心组件作用和介绍
    年薪30万+的HR这样做数据分析!(附关键指标&免费模版)
    X11 Xlib截屏问题及深入分析一 —— 源码位置
    C++系列十:C++函数
    set | map | multiset | multimap 快速上手
    逆序对专题
    人工智能和机器学习:走向智能未来的关键
  • 原文地址:https://blog.csdn.net/Cyril_KI/article/details/125479940
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号