mirror of
https://github.com/wassname/DeepTime.git
synced 2026-08-10 11:50:07 +08:00
wip
This commit is contained in:
+32
-19
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user