ccloud0525 commited on
Commit ·
c39e45e
1
Parent(s): b0d5f8a
feat: 'README'
Browse files- modeling_flame.py +16 -11
modeling_flame.py
CHANGED
|
@@ -550,8 +550,8 @@ class TimeDelayEmbedding(nn.Module):
|
|
| 550 |
embedding = rearrange(embedding, '(b n) d -> b n d', b=b, n=n)
|
| 551 |
return embedding
|
| 552 |
|
| 553 |
-
def forward(self, x, period=None
|
| 554 |
-
if not self.training
|
| 555 |
return self._embedding(x, period=period)
|
| 556 |
|
| 557 |
if period is None:
|
|
@@ -604,14 +604,14 @@ class MixedEmbedding(nn.Module):
|
|
| 604 |
proj.weight.data = resample(old=self.patch_embedding.proj.weight.data, new_patch_len=patch_len)
|
| 605 |
|
| 606 |
patch_embedding = proj(patches)
|
| 607 |
-
time_delay_embedding = self.time_delay_embedding(x, period=patch_len
|
| 608 |
embedding = patch_embedding + time_delay_embedding
|
| 609 |
return embedding
|
| 610 |
|
| 611 |
-
def forward(self, x, inference_patch_len=48
|
| 612 |
# do patching
|
| 613 |
# padding for the original stride
|
| 614 |
-
if not self.training
|
| 615 |
return self._flex_embedding(x, inference_patch_len)
|
| 616 |
|
| 617 |
seq_len = x.shape[-1]
|
|
@@ -627,7 +627,7 @@ class MixedEmbedding(nn.Module):
|
|
| 627 |
patch_embedding = self.patch_embedding(patches)
|
| 628 |
|
| 629 |
# [batch_size, patch_num, d_model]
|
| 630 |
-
time_delay_embedding = self.time_delay_embedding(x
|
| 631 |
|
| 632 |
embedding = patch_embedding + time_delay_embedding
|
| 633 |
|
|
@@ -720,9 +720,9 @@ class FLAMEModel(nn.Module):
|
|
| 720 |
|
| 721 |
def _predict(self, input, pred_len, inference_patch_len=48, num_samples=None):
|
| 722 |
if num_samples is not None and num_samples > 1:
|
| 723 |
-
return self._prob_predict(input, pred_len, num_samples)
|
| 724 |
|
| 725 |
-
x_enc = self.embedding(input, inference_patch_len
|
| 726 |
|
| 727 |
x_rec = self.encoder(x_enc)
|
| 728 |
|
|
@@ -741,13 +741,13 @@ class FLAMEModel(nn.Module):
|
|
| 741 |
point_forecasts = rearrange(dec_out, 'b n p -> b (n p)')
|
| 742 |
return point_forecasts[:, :pred_len]
|
| 743 |
|
| 744 |
-
def _prob_predict(self, input, pred_len, num_samples=None):
|
| 745 |
|
| 746 |
-
x_enc = self.embedding(input,
|
| 747 |
|
| 748 |
x_rec = self.encoder(x_enc)
|
| 749 |
|
| 750 |
-
predict_token_num = math.ceil(pred_len /
|
| 751 |
weights = self._get_weights(predict_token_num).unsqueeze(0).unsqueeze(-1).to(input.device)
|
| 752 |
last_token = x_rec[:, -1:, :]
|
| 753 |
x_enc = weights * last_token.repeat(1, predict_token_num, 1)
|
|
@@ -757,6 +757,11 @@ class FLAMEModel(nn.Module):
|
|
| 757 |
dist = self._prob_head(x_dec)
|
| 758 |
|
| 759 |
samples = dist.sample((num_samples,))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 760 |
samples = rearrange(samples, 's (b n) p -> b s (n p) ', n=predict_token_num)[:, :, :pred_len]
|
| 761 |
|
| 762 |
prob_forecasts = samples
|
|
|
|
| 550 |
embedding = rearrange(embedding, '(b n) d -> b n d', b=b, n=n)
|
| 551 |
return embedding
|
| 552 |
|
| 553 |
+
def forward(self, x, period=None):
|
| 554 |
+
if not self.training:
|
| 555 |
return self._embedding(x, period=period)
|
| 556 |
|
| 557 |
if period is None:
|
|
|
|
| 604 |
proj.weight.data = resample(old=self.patch_embedding.proj.weight.data, new_patch_len=patch_len)
|
| 605 |
|
| 606 |
patch_embedding = proj(patches)
|
| 607 |
+
time_delay_embedding = self.time_delay_embedding(x, period=patch_len)
|
| 608 |
embedding = patch_embedding + time_delay_embedding
|
| 609 |
return embedding
|
| 610 |
|
| 611 |
+
def forward(self, x, inference_patch_len=48):
|
| 612 |
# do patching
|
| 613 |
# padding for the original stride
|
| 614 |
+
if not self.training:
|
| 615 |
return self._flex_embedding(x, inference_patch_len)
|
| 616 |
|
| 617 |
seq_len = x.shape[-1]
|
|
|
|
| 627 |
patch_embedding = self.patch_embedding(patches)
|
| 628 |
|
| 629 |
# [batch_size, patch_num, d_model]
|
| 630 |
+
time_delay_embedding = self.time_delay_embedding(x)
|
| 631 |
|
| 632 |
embedding = patch_embedding + time_delay_embedding
|
| 633 |
|
|
|
|
| 720 |
|
| 721 |
def _predict(self, input, pred_len, inference_patch_len=48, num_samples=None):
|
| 722 |
if num_samples is not None and num_samples > 1:
|
| 723 |
+
return self._prob_predict(input, pred_len, num_samples, inference_patch_len)
|
| 724 |
|
| 725 |
+
x_enc = self.embedding(input, inference_patch_len)
|
| 726 |
|
| 727 |
x_rec = self.encoder(x_enc)
|
| 728 |
|
|
|
|
| 741 |
point_forecasts = rearrange(dec_out, 'b n p -> b (n p)')
|
| 742 |
return point_forecasts[:, :pred_len]
|
| 743 |
|
| 744 |
+
def _prob_predict(self, input, pred_len, num_samples=None, inference_patch_len=48):
|
| 745 |
|
| 746 |
+
x_enc = self.embedding(input, inference_patch_len=inference_patch_len)
|
| 747 |
|
| 748 |
x_rec = self.encoder(x_enc)
|
| 749 |
|
| 750 |
+
predict_token_num = math.ceil(pred_len / inference_patch_len)
|
| 751 |
weights = self._get_weights(predict_token_num).unsqueeze(0).unsqueeze(-1).to(input.device)
|
| 752 |
last_token = x_rec[:, -1:, :]
|
| 753 |
x_enc = weights * last_token.repeat(1, predict_token_num, 1)
|
|
|
|
| 757 |
dist = self._prob_head(x_dec)
|
| 758 |
|
| 759 |
samples = dist.sample((num_samples,))
|
| 760 |
+
|
| 761 |
+
weights = torch.eye(self.patch_len, device=x_dec.device)
|
| 762 |
+
resampled_weights = resample(old=weights, new_patch_len=inference_patch_len).T
|
| 763 |
+
|
| 764 |
+
samples = F.linear(samples, resampled_weights)
|
| 765 |
samples = rearrange(samples, 's (b n) p -> b s (n p) ', n=predict_token_num)[:, :, :pred_len]
|
| 766 |
|
| 767 |
prob_forecasts = samples
|