ccloud0525 commited on
Commit
c39e45e
·
1 Parent(s): b0d5f8a

feat: 'README'

Browse files
Files changed (1) hide show
  1. 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, num_samples=1):
554
- if not self.training and num_samples == 1:
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, num_samples=1)
608
  embedding = patch_embedding + time_delay_embedding
609
  return embedding
610
 
611
- def forward(self, x, inference_patch_len=48, num_samples=1):
612
  # do patching
613
  # padding for the original stride
614
- if not self.training and num_samples == 1:
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, num_samples=num_samples)
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, num_samples=1)
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, num_samples=num_samples)
747
 
748
  x_rec = self.encoder(x_enc)
749
 
750
- predict_token_num = math.ceil(pred_len / self.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,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