sparsetrace commited on
Commit
df6c670
·
verified ·
1 Parent(s): 47d2dad

Update PTST.py

Browse files
Files changed (1) hide show
  1. PTST.py +10 -7
PTST.py CHANGED
@@ -504,20 +504,23 @@ class PTST:
504
  ctx_n = (ctx - self._mu) / self._sd # (L_eff, D)
505
 
506
  preds_n_all = []
507
-
508
- # iterative chunking if steps > pred_len
509
- while len(preds_n_all) < H:
510
  ctx_n_j = jnp.asarray(ctx_n[None, :, :]) # (1,L_eff,D)
511
  block_n = np.array(self._predict_block_norm(ctx_n_j)[0]) # (pred_len, D)
512
-
513
- remaining = H - len(preds_n_all)
514
  take = min(self.pred_len, remaining)
 
515
  preds_n_all.append(block_n[:take])
516
-
517
- # roll context: append predictions, keep last L_eff
 
518
  ctx_n = np.concatenate([ctx_n, block_n[:take]], axis=0)
519
  ctx_n = ctx_n[-self.L_eff:, :]
520
 
 
521
  preds_n = np.concatenate(preds_n_all, axis=0) # (H, D)
522
  preds = preds_n * self._sd + self._mu
523
  return preds
 
504
  ctx_n = (ctx - self._mu) / self._sd # (L_eff, D)
505
 
506
  preds_n_all = []
507
+ n_done = 0
508
+
509
+ while n_done < H:
510
  ctx_n_j = jnp.asarray(ctx_n[None, :, :]) # (1,L_eff,D)
511
  block_n = np.array(self._predict_block_norm(ctx_n_j)[0]) # (pred_len, D)
512
+
513
+ remaining = H - n_done
514
  take = min(self.pred_len, remaining)
515
+
516
  preds_n_all.append(block_n[:take])
517
+ n_done += take
518
+
519
+ # roll context
520
  ctx_n = np.concatenate([ctx_n, block_n[:take]], axis=0)
521
  ctx_n = ctx_n[-self.L_eff:, :]
522
 
523
+
524
  preds_n = np.concatenate(preds_n_all, axis=0) # (H, D)
525
  preds = preds_n * self._sd + self._mu
526
  return preds