| 75 | """ |
| 76 | |
| 77 | def __init__(self, configs): |
| 78 | super(MultiScaleTrendMixing, self).__init__() |
| 79 | |
| 80 | self.up_sampling_layers = torch.nn.ModuleList( |
| 81 | [ |
| 82 | nn.Sequential( |
| 83 | torch.nn.Linear( |
| 84 | configs.seq_len // (configs.down_sampling_window ** (i + 1)), |
| 85 | configs.seq_len // (configs.down_sampling_window ** i), |
| 86 | ), |
| 87 | nn.GELU(), |
| 88 | torch.nn.Linear( |
| 89 | configs.seq_len // (configs.down_sampling_window ** i), |
| 90 | configs.seq_len // (configs.down_sampling_window ** i), |
| 91 | ), |
| 92 | ) |
| 93 | for i in reversed(range(configs.down_sampling_layers)) |
| 94 | ]) |
| 95 | |
| 96 | def forward(self, trend_list): |
| 97 | |