Update PTST.py
Browse files
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 |
-
|
| 509 |
-
while
|
| 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 -
|
| 514 |
take = min(self.pred_len, remaining)
|
|
|
|
| 515 |
preds_n_all.append(block_n[:take])
|
| 516 |
-
|
| 517 |
-
|
|
|
|
| 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
|