| 162 | return self.dropout(x) |
| 163 | |
| 164 | class DataEmbedding_mine(nn.Module): |
| 165 | def __init__(self, c_in, d_model, embed_type='fixed', freq='h', dropout=0.1, is_decoder=False): |
| 166 | super(DataEmbedding_mine, self).__init__() |
| 167 | if is_decoder: |
| 168 | c_in += 1 |
| 169 | self.value_embedding = TokenEmbedding(c_in=c_in, d_model=d_model) |
| 170 | self.position_embedding = PositionalEmbedding_new(d_model=d_model) |
| 171 | self.temporal_embedding = TimeFeatureEmbedding_new(d_model=d_model, embed_type=embed_type, freq=freq) |
| 172 | self.dropout = nn.Dropout(p=dropout) |
| 173 | self.is_decoder = is_decoder |
| 174 | |
| 175 | def forward(self, x, x_mark, scale, first_scale, label_len): |
| 176 | if self.is_decoder: |
| 177 | x = torch.cat((x, torch.ones((x.shape[0], x.shape[1], 1), device=x.device)), dim=2) |
| 178 | if scale==first_scale: |
| 179 | x[:,:label_len//scale,-1] = 0 |
| 180 | x[:,label_len//scale:,-1] = 0.5 |
| 181 | else: |
| 182 | x[:,:label_len//scale,-1] = 0 |
| 183 | x[:,label_len//scale:,-1] = 1 |
| 184 | vembed = self.value_embedding(x) |
| 185 | pembed = self.position_embedding(x, scale) |
| 186 | tembed = self.temporal_embedding(x_mark, scale) |
| 187 | x = vembed + pembed + tembed |
| 188 | return self.dropout(x) |
| 189 | |
| 190 | |
| 191 | class DataEmbedding_wo_pos(nn.Module): |