| 138 | self.concats = dict() |
| 139 | |
| 140 | def forward(self, x, scale=1): |
| 141 | if (scale, x.shape[0], x.shape[1]) not in self.concats: |
| 142 | concat_tensor = torch.tensor([[[1/scale-0.5]]]).cuda().repeat(x.shape[0],x.shape[1],1) |
| 143 | self.concats[(scale, x.shape[0], x.shape[1])] = concat_tensor |
| 144 | else: |
| 145 | concat_tensor = self.concats[(scale, x.shape[0], x.shape[1])] |
| 146 | x = torch.cat((x, concat_tensor), 2) |
| 147 | return self.embed(x) |
| 148 | |
| 149 | class DataEmbedding(nn.Module): |
| 150 | def __init__(self, c_in, d_model, embed_type='fixed', freq='h', dropout=0.1): |