MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / TemporalEmbedding

Class TemporalEmbedding

layers/Embed.py:66–93  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

64
65
66class TemporalEmbedding(nn.Module):
67 def __init__(self, d_model, embed_type='fixed', freq='h'):
68 super(TemporalEmbedding, self).__init__()
69
70 minute_size = 4
71 hour_size = 24
72 weekday_size = 7
73 day_size = 32
74 month_size = 13
75
76 Embed = FixedEmbedding if embed_type == 'fixed' else nn.Embedding
77 if freq == 't':
78 self.minute_embed = Embed(minute_size, d_model)
79 self.hour_embed = Embed(hour_size, d_model)
80 self.weekday_embed = Embed(weekday_size, d_model)
81 self.day_embed = Embed(day_size, d_model)
82 self.month_embed = Embed(month_size, d_model)
83
84 def forward(self, x):
85 x = x.long()
86 minute_x = self.minute_embed(x[:, :, 4]) if hasattr(
87 self, 'minute_embed') else 0.
88 hour_x = self.hour_embed(x[:, :, 3])
89 weekday_x = self.weekday_embed(x[:, :, 2])
90 day_x = self.day_embed(x[:, :, 1])
91 month_x = self.month_embed(x[:, :, 0])
92
93 return hour_x + weekday_x + day_x + month_x + minute_x
94
95
96class TimeFeatureEmbedding(nn.Module):

Callers 3

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected