This commit is contained in:
Alex Blank
2025-05-19 13:59:16 +02:00
parent 426f4d6963
commit c6defa2065
196 changed files with 18625 additions and 1 deletions
+121
View File
@@ -0,0 +1,121 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from layers.Embed import DataEmbedding, DataEmbedding_wo_pos,DataEmbedding_wo_pos_temp,DataEmbedding_wo_temp
from layers.AutoCorrelation import AutoCorrelation, AutoCorrelationLayer
from layers.Autoformer_EncDec import Encoder, Decoder, EncoderLayer, DecoderLayer, my_Layernorm, series_decomp
import math
import numpy as np
class Model(nn.Module):
"""
Autoformer is the first method to achieve the series-wise connection,
with inherent O(LlogL) complexity
"""
def __init__(self, configs):
super(Model, self).__init__()
self.seq_len = configs.seq_len
self.label_len = configs.label_len
self.pred_len = configs.pred_len
self.output_attention = configs.output_attention
# Decomp
kernel_size = configs.moving_avg
self.decomp = series_decomp(kernel_size)
# Embedding
# The series-wise connection inherently contains the sequential information.
# Thus, we can discard the position embedding of transformers.
if configs.embed_type == 0:
self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 1:
self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 2:
self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 3:
self.enc_embedding = DataEmbedding_wo_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 4:
self.enc_embedding = DataEmbedding_wo_pos_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
# Encoder
self.encoder = Encoder(
[
EncoderLayer(
AutoCorrelationLayer(
AutoCorrelation(False, configs.factor, attention_dropout=configs.dropout,
output_attention=configs.output_attention),
configs.d_model, configs.n_heads),
configs.d_model,
configs.d_ff,
moving_avg=configs.moving_avg,
dropout=configs.dropout,
activation=configs.activation
) for l in range(configs.e_layers)
],
norm_layer=my_Layernorm(configs.d_model)
)
# Decoder
self.decoder = Decoder(
[
DecoderLayer(
AutoCorrelationLayer(
AutoCorrelation(True, configs.factor, attention_dropout=configs.dropout,
output_attention=False),
configs.d_model, configs.n_heads),
AutoCorrelationLayer(
AutoCorrelation(False, configs.factor, attention_dropout=configs.dropout,
output_attention=False),
configs.d_model, configs.n_heads),
configs.d_model,
configs.c_out,
configs.d_ff,
moving_avg=configs.moving_avg,
dropout=configs.dropout,
activation=configs.activation,
)
for l in range(configs.d_layers)
],
norm_layer=my_Layernorm(configs.d_model),
projection=nn.Linear(configs.d_model, configs.c_out, bias=True)
)
def forward(self, x_enc, x_mark_enc, x_dec, x_mark_dec,
enc_self_mask=None, dec_self_mask=None, dec_enc_mask=None):
# decomp init
mean = torch.mean(x_enc, dim=1).unsqueeze(1).repeat(1, self.pred_len, 1)
zeros = torch.zeros([x_dec.shape[0], self.pred_len, x_dec.shape[2]], device=x_enc.device)
seasonal_init, trend_init = self.decomp(x_enc)
# decoder input
trend_init = torch.cat([trend_init[:, -self.label_len:, :], mean], dim=1)
seasonal_init = torch.cat([seasonal_init[:, -self.label_len:, :], zeros], dim=1)
# enc
enc_out = self.enc_embedding(x_enc, x_mark_enc)
enc_out, attns = self.encoder(enc_out, attn_mask=enc_self_mask)
# dec
dec_out = self.dec_embedding(seasonal_init, x_mark_dec)
seasonal_part, trend_part = self.decoder(dec_out, enc_out, x_mask=dec_self_mask, cross_mask=dec_enc_mask,
trend=trend_init)
# final
dec_out = trend_part + seasonal_part
if self.output_attention:
return dec_out[:, -self.pred_len:, :], attns
else:
return dec_out[:, -self.pred_len:, :] # [B, L, D]
+87
View File
@@ -0,0 +1,87 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class moving_avg(nn.Module):
"""
Moving average block to highlight the trend of time series
"""
def __init__(self, kernel_size, stride):
super(moving_avg, self).__init__()
self.kernel_size = kernel_size
self.avg = nn.AvgPool1d(kernel_size=kernel_size, stride=stride, padding=0)
def forward(self, x):
# padding on the both ends of time series
front = x[:, 0:1, :].repeat(1, (self.kernel_size - 1) // 2, 1)
end = x[:, -1:, :].repeat(1, (self.kernel_size - 1) // 2, 1)
x = torch.cat([front, x, end], dim=1)
x = self.avg(x.permute(0, 2, 1))
x = x.permute(0, 2, 1)
return x
class series_decomp(nn.Module):
"""
Series decomposition block
"""
def __init__(self, kernel_size):
super(series_decomp, self).__init__()
self.moving_avg = moving_avg(kernel_size, stride=1)
def forward(self, x):
moving_mean = self.moving_avg(x)
res = x - moving_mean
return res, moving_mean
class Model(nn.Module):
"""
Decomposition-Linear
"""
def __init__(self, configs):
super(Model, self).__init__()
self.seq_len = configs.seq_len
self.pred_len = configs.pred_len
# Decompsition Kernel Size
kernel_size = 25
self.decompsition = series_decomp(kernel_size)
self.individual = configs.individual
self.channels = configs.enc_in
if self.individual:
self.Linear_Seasonal = nn.ModuleList()
self.Linear_Trend = nn.ModuleList()
for i in range(self.channels):
self.Linear_Seasonal.append(nn.Linear(self.seq_len,self.pred_len))
self.Linear_Trend.append(nn.Linear(self.seq_len,self.pred_len))
# Use this two lines if you want to visualize the weights
# self.Linear_Seasonal[i].weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
# self.Linear_Trend[i].weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
else:
self.Linear_Seasonal = nn.Linear(self.seq_len,self.pred_len)
self.Linear_Trend = nn.Linear(self.seq_len,self.pred_len)
# Use this two lines if you want to visualize the weights
# self.Linear_Seasonal.weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
# self.Linear_Trend.weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
def forward(self, x):
# x: [Batch, Input length, Channel]
seasonal_init, trend_init = self.decompsition(x)
seasonal_init, trend_init = seasonal_init.permute(0,2,1), trend_init.permute(0,2,1)
if self.individual:
seasonal_output = torch.zeros([seasonal_init.size(0),seasonal_init.size(1),self.pred_len],dtype=seasonal_init.dtype).to(seasonal_init.device)
trend_output = torch.zeros([trend_init.size(0),trend_init.size(1),self.pred_len],dtype=trend_init.dtype).to(trend_init.device)
for i in range(self.channels):
seasonal_output[:,i,:] = self.Linear_Seasonal[i](seasonal_init[:,i,:])
trend_output[:,i,:] = self.Linear_Trend[i](trend_init[:,i,:])
else:
seasonal_output = self.Linear_Seasonal(seasonal_init)
trend_output = self.Linear_Trend(trend_init)
x = seasonal_output + trend_output
return x.permute(0,2,1) # to [Batch, Output length, Channel]
+101
View File
@@ -0,0 +1,101 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from utils.masking import TriangularCausalMask, ProbMask
from layers.Transformer_EncDec import Decoder, DecoderLayer, Encoder, EncoderLayer, ConvLayer
from layers.SelfAttention_Family import FullAttention, ProbAttention, AttentionLayer
from layers.Embed import DataEmbedding,DataEmbedding_wo_pos,DataEmbedding_wo_temp,DataEmbedding_wo_pos_temp
import numpy as np
class Model(nn.Module):
"""
Informer with Propspare attention in O(LlogL) complexity
"""
def __init__(self, configs):
super(Model, self).__init__()
self.pred_len = configs.pred_len
self.output_attention = configs.output_attention
# Embedding
if configs.embed_type == 0:
self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 1:
self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 2:
self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 3:
self.enc_embedding = DataEmbedding_wo_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 4:
self.enc_embedding = DataEmbedding_wo_pos_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
# Encoder
self.encoder = Encoder(
[
EncoderLayer(
AttentionLayer(
ProbAttention(False, configs.factor, attention_dropout=configs.dropout,
output_attention=configs.output_attention),
configs.d_model, configs.n_heads),
configs.d_model,
configs.d_ff,
dropout=configs.dropout,
activation=configs.activation
) for l in range(configs.e_layers)
],
[
ConvLayer(
configs.d_model
) for l in range(configs.e_layers - 1)
] if configs.distil else None,
norm_layer=torch.nn.LayerNorm(configs.d_model)
)
# Decoder
self.decoder = Decoder(
[
DecoderLayer(
AttentionLayer(
ProbAttention(True, configs.factor, attention_dropout=configs.dropout, output_attention=False),
configs.d_model, configs.n_heads),
AttentionLayer(
ProbAttention(False, configs.factor, attention_dropout=configs.dropout, output_attention=False),
configs.d_model, configs.n_heads),
configs.d_model,
configs.d_ff,
dropout=configs.dropout,
activation=configs.activation,
)
for l in range(configs.d_layers)
],
norm_layer=torch.nn.LayerNorm(configs.d_model),
projection=nn.Linear(configs.d_model, configs.c_out, bias=True)
)
def forward(self, x_enc, x_mark_enc, x_dec, x_mark_dec,
enc_self_mask=None, dec_self_mask=None, dec_enc_mask=None):
enc_out = self.enc_embedding(x_enc, x_mark_enc)
enc_out, attns = self.encoder(enc_out, attn_mask=enc_self_mask)
dec_out = self.dec_embedding(x_dec, x_mark_dec)
dec_out = self.decoder(dec_out, enc_out, x_mask=dec_self_mask, cross_mask=dec_enc_mask)
if self.output_attention:
return dec_out[:, -self.pred_len:, :], attns
else:
return dec_out[:, -self.pred_len:, :] # [B, L, D]
+21
View File
@@ -0,0 +1,21 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class Model(nn.Module):
"""
Just one Linear layer
"""
def __init__(self, configs):
super(Model, self).__init__()
self.seq_len = configs.seq_len
self.pred_len = configs.pred_len
self.Linear = nn.Linear(self.seq_len, self.pred_len)
# Use this line if you want to visualize the weights
# self.Linear.weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
def forward(self, x):
# x: [Batch, Input length, Channel]
x = self.Linear(x.permute(0,2,1)).permute(0,2,1)
return x # [Batch, Output length, Channel]
+24
View File
@@ -0,0 +1,24 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class Model(nn.Module):
"""
Normalization-Linear
"""
def __init__(self, configs):
super(Model, self).__init__()
self.seq_len = configs.seq_len
self.pred_len = configs.pred_len
self.Linear = nn.Linear(self.seq_len, self.pred_len)
# Use this line if you want to visualize the weights
# self.Linear.weight = nn.Parameter((1/self.seq_len)*torch.ones([self.pred_len,self.seq_len]))
def forward(self, x):
# x: [Batch, Input length, Channel]
seq_last = x[:,-1:,:].detach()
x = x - seq_last
x = self.Linear(x.permute(0,2,1)).permute(0,2,1)
x = x + seq_last
return x # [Batch, Output length, Channel]
+127
View File
@@ -0,0 +1,127 @@
__all__ = ['PatchTST']
# Cell
from typing import Callable, Optional
import torch
from torch import nn
from torch import Tensor
import torch.nn.functional as F
import numpy as np
from models.third_party.patch_tst.layers.PatchTST_backbone import PatchTST_backbone
from models.third_party.patch_tst.layers.PatchTST_layers import series_decomp
class Model(nn.Module):
def __init__(self, input_dim: int, output_dim: int, configs, max_seq_len: Optional[int] = 1024,
d_k: Optional[int] = None, d_v: Optional[int] = None,
norm: str = 'BatchNorm', attn_dropout: float = 0.,
act: str = "gelu", key_padding_mask: bool = 'auto', padding_var: Optional[int] = None,
attn_mask: Optional[Tensor] = None, res_attention: bool = True,
pre_norm: bool = False, store_attn: bool = False, pe: str = 'zeros', learn_pe: bool = True,
pretrain_head: bool = False, head_type='flatten', verbose: bool = False, **kwargs):
super().__init__()
# load parameters
c_in = input_dim
context_window = configs['seq_len']
target_window = configs['pred_len']
dec_out = output_dim
seq_pred = configs["seq_pred"]
n_layers = configs['e_layers']
n_heads = configs['n_heads']
d_model = configs['d_model']
d_ff = configs['d_ff']
dropout = configs['dropout']
fc_dropout = configs['fc_dropout']
head_dropout = configs['head_dropout']
individual = configs['individual']
patch_len = configs['patch_len']
stride = configs['stride']
padding_patch = configs['padding_patch']
revin = configs['revin']
affine = configs['affine']
subtract_last = configs['subtract_last']
decomposition = configs['decomposition']
kernel_size = configs['kernel_size']
# model
self.decomposition = decomposition
if self.decomposition:
self.decomp_module = series_decomp(kernel_size)
self.model_trend = PatchTST_backbone(c_in=c_in, context_window=context_window, target_window=target_window,
# extras
dec_out=dec_out,
seq_pred=seq_pred,
#
patch_len=patch_len, stride=stride,
max_seq_len=max_seq_len, n_layers=n_layers, d_model=d_model,
n_heads=n_heads, d_k=d_k, d_v=d_v, d_ff=d_ff, norm=norm,
attn_dropout=attn_dropout,
dropout=dropout, act=act, key_padding_mask=key_padding_mask,
padding_var=padding_var,
attn_mask=attn_mask, res_attention=res_attention, pre_norm=pre_norm,
store_attn=store_attn,
pe=pe, learn_pe=learn_pe, fc_dropout=fc_dropout,
head_dropout=head_dropout, padding_patch=padding_patch,
pretrain_head=pretrain_head, head_type=head_type,
individual=individual, revin=revin, affine=affine,
subtract_last=subtract_last, verbose=verbose, **kwargs)
self.model_res = PatchTST_backbone(c_in=c_in, context_window=context_window, target_window=target_window,
# extras
dec_out=dec_out,
seq_pred=seq_pred,
#
patch_len=patch_len, stride=stride,
max_seq_len=max_seq_len, n_layers=n_layers, d_model=d_model,
n_heads=n_heads, d_k=d_k, d_v=d_v, d_ff=d_ff, norm=norm,
attn_dropout=attn_dropout,
dropout=dropout, act=act, key_padding_mask=key_padding_mask,
padding_var=padding_var,
attn_mask=attn_mask, res_attention=res_attention, pre_norm=pre_norm,
store_attn=store_attn,
pe=pe, learn_pe=learn_pe, fc_dropout=fc_dropout,
head_dropout=head_dropout, padding_patch=padding_patch,
pretrain_head=pretrain_head, head_type=head_type, individual=individual,
revin=revin, affine=affine,
subtract_last=subtract_last, verbose=verbose, **kwargs)
else:
self.model = PatchTST_backbone(c_in=c_in, context_window=context_window, target_window=target_window,
# extras
dec_out=dec_out,
seq_pred=seq_pred,
#
patch_len=patch_len, stride=stride,
max_seq_len=max_seq_len, n_layers=n_layers, d_model=d_model,
n_heads=n_heads, d_k=d_k, d_v=d_v, d_ff=d_ff, norm=norm,
attn_dropout=attn_dropout,
dropout=dropout, act=act, key_padding_mask=key_padding_mask,
padding_var=padding_var,
attn_mask=attn_mask, res_attention=res_attention, pre_norm=pre_norm,
store_attn=store_attn,
pe=pe, learn_pe=learn_pe, fc_dropout=fc_dropout, head_dropout=head_dropout,
padding_patch=padding_patch,
pretrain_head=pretrain_head, head_type=head_type, individual=individual,
revin=revin, affine=affine,
subtract_last=subtract_last, verbose=verbose, **kwargs)
def forward(self, x): # x: [Batch, Input length, Channel]
if self.decomposition:
res_init, trend_init = self.decomp_module(x)
res_init, trend_init = res_init.permute(0, 2, 1), trend_init.permute(0, 2,
1) # x: [Batch, Channel, Input length]
res = self.model_res(res_init)
trend = self.model_trend(trend_init)
x = res + trend
x = x.permute(0, 2, 1) # x: [Batch, Input length, Channel]
else:
x = x.permute(0, 2, 1) # x: [Batch, Channel, Input length]
x = self.model(x)
x = x.permute(0, 2, 1) # x: [Batch, Input length, Channel]
return x
+120
View File
@@ -0,0 +1,120 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from tqdm import tqdm
import pmdarima as pm
import threading
from sklearn.ensemble import GradientBoostingRegressor
class Naive_repeat(nn.Module):
def __init__(self, configs):
super(Naive_repeat, self).__init__()
self.pred_len = configs.pred_len
def forward(self, x):
B,L,D = x.shape
x = x[:,-1,:].reshape(B,1,D).repeat(self.pred_len,axis=1)
return x # [B, L, D]
class Naive_thread(threading.Thread):
def __init__(self,func,args=()):
super(Naive_thread,self).__init__()
self.func = func
self.args = args
def run(self):
self.results = self.func(*self.args)
def return_result(self):
threading.Thread.join(self)
return self.results
def _arima(seq,pred_len,bt,i):
model = pm.auto_arima(seq)
forecasts = model.predict(pred_len)
return forecasts,bt,i
class Arima(nn.Module):
"""
Extremely slow, please sample < 0.1
"""
def __init__(self, configs):
super(Arima, self).__init__()
self.pred_len = configs.pred_len
def forward(self, x):
result = np.zeros([x.shape[0],self.pred_len,x.shape[2]])
threads = []
for bt,seqs in tqdm(enumerate(x)):
for i in range(seqs.shape[-1]):
seq = seqs[:,i]
one_seq = Naive_thread(func=_arima,args=(seq,self.pred_len,bt,i))
threads.append(one_seq)
threads[-1].start()
for every_thread in tqdm(threads):
forcast,bt,i = every_thread.return_result()
result[bt,:,i] = forcast
return result # [B, L, D]
def _sarima(season,seq,pred_len,bt,i):
model = pm.auto_arima(seq, seasonal=True, m=season)
forecasts = model.predict(pred_len)
return forecasts,bt,i
class SArima(nn.Module):
"""
Extremely extremely slow, please sample < 0.01
"""
def __init__(self, configs):
super(SArima, self).__init__()
self.pred_len = configs.pred_len
self.seq_len = configs.seq_len
self.season = 24
if 'Ettm' in configs.data_path:
self.season = 12
elif 'ILI' in configs.data_path:
self.season = 1
if self.season >= self.seq_len:
self.season = 1
def forward(self, x):
result = np.zeros([x.shape[0],self.pred_len,x.shape[2]])
threads = []
for bt,seqs in tqdm(enumerate(x)):
for i in range(seqs.shape[-1]):
seq = seqs[:,i]
one_seq = Naive_thread(func=_sarima,args=(self.season,seq,self.pred_len,bt,i))
threads.append(one_seq)
threads[-1].start()
for every_thread in tqdm(threads):
forcast,bt,i = every_thread.return_result()
result[bt,:,i] = forcast
return result # [B, L, D]
def _gbrt(seq,seq_len,pred_len,bt,i):
model = GradientBoostingRegressor()
model.fit(np.arange(seq_len).reshape(-1,1),seq.reshape(-1,1))
forecasts = model.predict(np.arange(seq_len,seq_len+pred_len).reshape(-1,1))
return forecasts,bt,i
class GBRT(nn.Module):
def __init__(self, configs):
super(GBRT, self).__init__()
self.seq_len = configs.seq_len
self.pred_len = configs.pred_len
def forward(self, x):
result = np.zeros([x.shape[0],self.pred_len,x.shape[2]])
threads = []
for bt,seqs in tqdm(enumerate(x)):
for i in range(seqs.shape[-1]):
seq = seqs[:,i]
one_seq = Naive_thread(func=_gbrt,args=(seq,self.seq_len,self.pred_len,bt,i))
threads.append(one_seq)
threads[-1].start()
for every_thread in tqdm(threads):
forcast,bt,i = every_thread.return_result()
result[bt,:,i] = forcast
return result # [B, L, D]
+94
View File
@@ -0,0 +1,94 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from layers.Transformer_EncDec import Decoder, DecoderLayer, Encoder, EncoderLayer, ConvLayer
from layers.SelfAttention_Family import FullAttention, AttentionLayer
from layers.Embed import DataEmbedding,DataEmbedding_wo_pos,DataEmbedding_wo_temp,DataEmbedding_wo_pos_temp
import numpy as np
class Model(nn.Module):
"""
Vanilla Transformer with O(L^2) complexity
"""
def __init__(self, configs):
super(Model, self).__init__()
self.pred_len = configs.pred_len
self.output_attention = configs.output_attention
# Embedding
if configs.embed_type == 0:
self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 1:
self.enc_embedding = DataEmbedding(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 2:
self.enc_embedding = DataEmbedding_wo_pos(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 3:
self.enc_embedding = DataEmbedding_wo_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
elif configs.embed_type == 4:
self.enc_embedding = DataEmbedding_wo_pos_temp(configs.enc_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
self.dec_embedding = DataEmbedding_wo_pos_temp(configs.dec_in, configs.d_model, configs.embed, configs.freq,
configs.dropout)
# Encoder
self.encoder = Encoder(
[
EncoderLayer(
AttentionLayer(
FullAttention(False, configs.factor, attention_dropout=configs.dropout,
output_attention=configs.output_attention), configs.d_model, configs.n_heads),
configs.d_model,
configs.d_ff,
dropout=configs.dropout,
activation=configs.activation
) for l in range(configs.e_layers)
],
norm_layer=torch.nn.LayerNorm(configs.d_model)
)
# Decoder
self.decoder = Decoder(
[
DecoderLayer(
AttentionLayer(
FullAttention(True, configs.factor, attention_dropout=configs.dropout, output_attention=False),
configs.d_model, configs.n_heads),
AttentionLayer(
FullAttention(False, configs.factor, attention_dropout=configs.dropout, output_attention=False),
configs.d_model, configs.n_heads),
configs.d_model,
configs.d_ff,
dropout=configs.dropout,
activation=configs.activation,
)
for l in range(configs.d_layers)
],
norm_layer=torch.nn.LayerNorm(configs.d_model),
projection=nn.Linear(configs.d_model, configs.c_out, bias=True)
)
def forward(self, x_enc, x_mark_enc, x_dec, x_mark_dec,
enc_self_mask=None, dec_self_mask=None, dec_enc_mask=None):
enc_out = self.enc_embedding(x_enc, x_mark_enc)
enc_out, attns = self.encoder(enc_out, attn_mask=enc_self_mask)
dec_out = self.dec_embedding(x_dec, x_mark_dec)
dec_out = self.decoder(dec_out, enc_out, x_mask=dec_self_mask, cross_mask=dec_enc_mask)
if self.output_attention:
return dec_out[:, -self.pred_len:, :], attns
else:
return dec_out[:, -self.pred_len:, :] # [B, L, D]