This commit is contained in:
wassname
2022-11-22 21:54:58 +08:00
parent 8818a2b17b
commit fc1a605b01
48 changed files with 803 additions and 1524 deletions
+32 -19
View File
@@ -11,40 +11,54 @@ import torch.nn as nn
from torch import Tensor
from einops import rearrange, repeat, reduce
from models.modules.causalinception import CausalInceptionTimePlus, CausalConv1d
from models.modules.metareghead import RegressionHead
from models.modules.causalinception import CausalInceptionTimePlus
from models.modules.inrplus2 import INRPlus2
from models.modules.regressors import RidgeRegressor
from models.modules.inr import INR
from models.modules.encoders import LSTMEncoder, TransformerEncoder2, TransformerEncoder, InceptionEncoder, LSTMEncoder2, MLPEncoder
# from models.modules.regressors import RidgeRegressor
@gin.configurable()
def deeptime3(dim_size:int, datetime_feats: int, layer_size: int, inr_layers: int, n_fourier_feats: int, scales: float):
return DeepTIMe3(dim_size, datetime_feats, layer_size, inr_layers, n_fourier_feats, scales)
def deeptime3(dim_size:int, datetime_feats: int, layer_size: int, inr_layers: int, n_fourier_feats: int, scales: float, dropout: float, base_learner: str, encoder:str, inr: str):
return DeepTIMe3(dim_size, datetime_feats, layer_size, inr_layers, n_fourier_feats, scales, dropout, base_learner, encoder, inr)
class DeepTIMe3(nn.Module):
def __init__(self, dim_size: int, datetime_feats: int, layer_size: int, inr_layers: int, n_fourier_feats: int, scales: float, dropout: float=0.3):
def __init__(self, dim_size: int, datetime_feats: int, layer_size: int, inr_layers: int, n_fourier_feats: int, scales: float, dropout: float=0.3, base_learner:str='Ridge', encoder:str='inception', inr:str='INR'):
super().__init__()
# encode the past
encoded_size = layer_size//2
self.encoder = CausalInceptionTimePlus(
c_in=dim_size, c_out=encoded_size,
# nf=32, depth=6,
nf=17, depth=3,
bn=True,
dilation=6,
ks=[39, 19, 3],
coord=True, fc_dropout=dropout,
)
if encoder == 'inception':
encoded_size = layer_size
self.encoder = CausalInceptionTimePlus(
c_in=dim_size, c_out=encoded_size,
# nf=24, depth=4,
nf=17, depth=3,
bn=True,
dilation=2,
ks=[39, 19, 3],
coord=True, fc_dropout=dropout,
)
elif encoder == 'lstm':
self.encoder = LSTMEncoder()
else:
raise NotADirectoryError(encoder)
# translate coords to a representation, given a summary of the past
coord_size = 1
in_feats=datetime_feats+encoded_size+coord_size
self.inr = INRPlus2(in_feats=in_feats, layers=inr_layers, layer_size=layer_size,
if inr=='INRPlus2':
self.inr = INRPlus2(in_feats=in_feats, layers=inr_layers, layer_size=layer_size,
n_fourier_feats=n_fourier_feats, scales=scales, dropout=dropout)
elif inr=="INR":
self.inr = INR(in_feats=in_feats, layers=inr_layers, layer_size=layer_size,
n_fourier_feats=n_fourier_feats, scales=scales, dropout=dropout)
else:
raise NotImplementedError(inr)
# meta learn y given a representation
self.adaptive_weights = RidgeRegressor()
self.regressionhead = RegressionHead(base_learner=base_learner, d=layer_size, dropout=dropout)
self.datetime_feats = datetime_feats
self.inr_layers = inr_layers
@@ -78,8 +92,7 @@ class DeepTIMe3(nn.Module):
context_reprs = self.encode_and_decode(context_past_x, context_time)
query_reprs = self.encode_and_decode(query_past_x, query_time, offset=context_reprs.shape[1])
w, b = self.adaptive_weights(context_reprs, context_y)
preds = self.forecast(query_reprs, w, b)
preds = self.regressionhead(query_reprs, context_reprs, context_y)
return preds
def forecast(self, inp: Tensor, w: Tensor, b: Tensor) -> Tensor: