| 107 | |
| 108 | |
| 109 | class DataEmbedding(nn.Module): |
| 110 | def __init__(self, c_in, d_model, embed_type='fixed', freq='h', dropout=0.1): |
| 111 | super(DataEmbedding, self).__init__() |
| 112 | self.c_in = c_in |
| 113 | self.d_model = d_model |
| 114 | self.value_embedding = TokenEmbedding(c_in=c_in, d_model=d_model) |
| 115 | self.position_embedding = PositionalEmbedding(d_model=d_model) |
| 116 | self.temporal_embedding = TemporalEmbedding(d_model=d_model, embed_type=embed_type, |
| 117 | freq=freq) if embed_type != 'timeF' else TimeFeatureEmbedding( |
| 118 | d_model=d_model, embed_type=embed_type, freq=freq) |
| 119 | self.dropout = nn.Dropout(p=dropout) |
| 120 | |
| 121 | def forward(self, x, x_mark): |
| 122 | _, _, N = x.size() |
| 123 | if N == self.c_in: |
| 124 | if x_mark is None: |
| 125 | x = self.value_embedding(x) + self.position_embedding(x) |
| 126 | else: |
| 127 | x = self.value_embedding( |
| 128 | x) + self.temporal_embedding(x_mark) + self.position_embedding(x) |
| 129 | elif N == self.d_model: |
| 130 | if x_mark is None: |
| 131 | x = x + self.position_embedding(x) |
| 132 | else: |
| 133 | x = x + self.temporal_embedding(x_mark) + self.position_embedding(x) |
| 134 | |
| 135 | return self.dropout(x) |
| 136 | |
| 137 | |
| 138 | class DataEmbedding_ms(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected