| 64 | |
| 65 | |
| 66 | class 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 | |
| 96 | class TimeFeatureEmbedding(nn.Module): |