Download forecast_wrapper.py from tensorlink-dev/yumoto-alpha-18m: direct link, hf CLI and curl.
- Browser
- Download file 17.3 kB
-
https://huggingface.co/tensorlink-dev/yumoto-alpha-18m/resolve/main/forecast_wrapper.py
- Command line
-
hf download hf://tensorlink-dev/yumoto-alpha-18m/forecast_wrapper.py
-
curl -L -o forecast_wrapper.py https://huggingface.co/tensorlink-dev/yumoto-alpha-18m/resolve/main/forecast_wrapper.py
17.3 kB
| """Auto-generated by cascade_model. Loads an arm checkpoint and decodes the | |
| full horizon in one forward pass via contiguous patch masking. | |
| forecast_quantiles_batch(histories, horizon) -> (B, horizon, num_q) | |
| forecast_quantiles(history, horizon) -> (1, horizon, num_q) | |
| forecast(history, horizon, num_samples) -> (1, num_samples, horizon) | |
| """ | |
| from __future__ import annotations | |
| import hashlib | |
| import importlib.util | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| STABLE_DECODE_STEPS = 768 | |
| def _hurdle_mix(res_row, ctx_row, levels=None): | |
| """Apply the zero-movement hurdle to one row's (..., H, Q) quantiles.""" | |
| lv = np.asarray(levels if levels is not None else | |
| [.1, .2, .3, .4, .5, .6, .7, .8, .9], dtype=np.float64) | |
| tail = ctx_row[-2016:] | |
| if len(tail) < 128: | |
| return res_row | |
| d = np.abs(np.diff(tail)) | |
| p_move = float(np.mean(d > 1e-12)) | |
| if p_move > 0.60: # live asset: exact no-op | |
| return res_row | |
| p_move = max(p_move, 1e-4) | |
| v = float(tail[-1]) | |
| Hn = res_row.shape[-2] | |
| out = res_row.copy() | |
| for t in range(Hn): | |
| pi = (1.0 - p_move) ** (t + 1) | |
| if pi < 1e-3: | |
| break # later steps: even less mass | |
| q = np.sort(res_row[..., t, :], axis=-1) | |
| flat_q = q.reshape(-1, q.shape[-1]) | |
| flat_o = out[..., t, :].reshape(-1, q.shape[-1]) | |
| for r in range(flat_q.shape[0]): | |
| qr = flat_q[r] | |
| Fv = float(np.interp(v, qr, lv, left=0.0, right=1.0)) | |
| lo = (1.0 - pi) * Fv | |
| hi = lo + pi | |
| newq = np.empty_like(qr) | |
| for j, l in enumerate(lv): | |
| if l < lo: | |
| newq[j] = float(np.interp(l / (1.0 - pi), lv, qr)) | |
| elif l <= hi: | |
| newq[j] = v | |
| else: | |
| newq[j] = float(np.interp((l - pi) / (1.0 - pi), lv, qr)) | |
| flat_o[r] = newq | |
| return out | |
| def _load_model_module(d: Path): | |
| spec = importlib.util.spec_from_file_location("cascade_model_ckpt", d / "model.py") | |
| mod = importlib.util.module_from_spec(spec) | |
| sys.modules[spec.name] = mod | |
| spec.loader.exec_module(mod) | |
| return mod | |
| class Wrapper: | |
| def __init__(self, checkpoint_dir, device: str = "cpu"): | |
| d = Path(checkpoint_dir) | |
| self.device = device | |
| cfg_obj = json.loads((d / "config.json").read_text()) | |
| self.m = _load_model_module(d) | |
| self.cfg = self.m.CascadeModelConfig(**cfg_obj["config"]) | |
| self.quantile_levels = [float(v) for v in cfg_obj["quantile_levels"]] | |
| self.levels = torch.tensor(self.quantile_levels, dtype=torch.float32, device=device) | |
| self.model = self.m.CascadeModel(self.cfg).to(device).eval() | |
| from safetensors.torch import load_file | |
| self.model.load_state_dict(load_file(str(d / "weights.safetensors"))) | |
| def _prep(self, histories): | |
| """Return (context, missing_mask), both (B, window_len). | |
| Real series have gaps. GIFT-Eval's electricity is 25% NaN, bitbrains | |
| 14%, car_parts 12% — and an unhandled NaN propagates through the causal | |
| scaler into every prediction, so the whole config fails with "Forecast | |
| contains NaN values". That silently cost 23 of 97 configs, concentrated | |
| on exactly the messy observability data this model is for. | |
| The architecture already handles this: the model takes a binary mask | |
| channel (1 = unobserved), and causal_standardize excludes masked entries | |
| from its statistics so they carry the last observed stats forward. It is | |
| the same mechanism training uses for CPM. Inference simply never | |
| populated it from NaNs. | |
| Masked positions are zero-filled in the value channel — the model sees | |
| the mask bit, not the filler. | |
| """ | |
| ps = self.cfg.patch_size | |
| n_ctx = max(2, self.cfg.context_length // ps) | |
| window = n_ctx * ps | |
| rows, masks = [], [] | |
| for h in histories: | |
| h = np.asarray(h, dtype=np.float64).reshape(-1) | |
| miss = ~np.isfinite(h) | |
| if h.shape[0] < window: | |
| # Left-pad with the first OBSERVED value, and mark the pad | |
| # unobserved so it cannot bias the causal statistics. | |
| obs = h[~miss] | |
| fill = obs[0] if obs.size else 0.0 | |
| pad_n = window - h.shape[0] | |
| h = np.concatenate([np.full(pad_n, fill), h]) | |
| miss = np.concatenate([np.ones(pad_n, dtype=bool), miss]) | |
| else: | |
| h, miss = h[-window:], miss[-window:] | |
| if miss.all(): | |
| # Nothing observed at all: fall back to zeros, all-observed, so | |
| # the scaler's eps floor keeps the forward pass finite. | |
| h = np.zeros_like(h); miss = np.zeros_like(miss) | |
| else: | |
| h = np.where(miss, 0.0, h) | |
| rows.append(h) | |
| masks.append(miss.astype(np.float64)) | |
| x = torch.as_tensor(np.stack(rows), dtype=torch.float64, device=self.device) | |
| m = torch.as_tensor(np.stack(masks), dtype=torch.float64, device=self.device) | |
| return x, m | |
| def _decode_block_z(self, z, block: int, miss=None): | |
| ps = self.cfg.patch_size | |
| ctx_p = min(z.shape[1] // ps, self.cfg.max_patches - block) | |
| ctx = z[:, -ctx_p * ps:].view(z.shape[0], ctx_p, ps) | |
| filler = torch.zeros(z.shape[0], block, ps, dtype=ctx.dtype, device=self.device) | |
| # Per-ENTRY mask: horizon patches are fully unobserved, and context | |
| # patches carry whichever entries were missing in the history. A | |
| # patch-level mask would be too coarse — one missing step would blank | |
| # the whole 32-step patch. | |
| m_ctx = (torch.zeros(z.shape[0], ctx_p * ps, dtype=ctx.dtype, device=self.device) | |
| if miss is None else miss[:, -ctx_p * ps:].to(ctx.dtype)) | |
| m_ctx = m_ctx.view(z.shape[0], ctx_p, ps) | |
| m_hz = torch.ones(z.shape[0], block, ps, dtype=ctx.dtype, device=self.device) | |
| mask = torch.cat([m_ctx, m_hz], dim=1) | |
| pred = self.model(torch.cat([ctx, filler], dim=1), mask=mask) | |
| q = pred[:, ctx_p - 1: ctx_p + block - 1] | |
| q, _ = torch.sort(q, dim=-1) | |
| return q.reshape(z.shape[0], block * ps, -1) | |
| def _decode_quantiles(self, x, horizon: int, miss=None): | |
| ps = self.cfg.patch_size | |
| stable = max(1, min(STABLE_DECODE_STEPS // ps, self.cfg.max_patches - 2)) | |
| remaining = -(-int(horizon) // ps) | |
| lo = hi = None | |
| out = [] | |
| if miss is None: | |
| miss = torch.zeros_like(x) | |
| while remaining > 0: | |
| block = min(remaining, stable) | |
| # Mask-aware scaling: missing entries are excluded from the causal | |
| # statistics, so they carry the last observed loc/scale forward | |
| # instead of poisoning them with NaN. | |
| z, loc_t, scale_t = self.m.causal_standardize( | |
| x, mask=miss, binary_passthrough=self.cfg.binary_passthrough | |
| ) | |
| loc = loc_t[:, -1:].double().unsqueeze(-1) | |
| scale = scale_t[:, -1:].double().unsqueeze(-1) | |
| if lo is None: | |
| # Clamp bounds from the OBSERVED range only — a masked entry is | |
| # zero-filled, and letting that zero set the bound would drag | |
| # the clamp toward the origin on a series that never visits it. | |
| obs = torch.where(miss > 0, torch.nan, x) | |
| lo = obs.nan_to_num(nan=float("inf")).min(dim=-1, keepdim=True).values.unsqueeze(-1) - 1e4 * scale | |
| hi = obs.nan_to_num(nan=float("-inf")).max(dim=-1, keepdim=True).values.unsqueeze(-1) + 1e4 * scale | |
| qz = self._decode_block_z(z.to(torch.float32), block, miss=miss) | |
| q = torch.sinh(qz.double()) * scale + loc | |
| q = torch.clamp(q, min=lo, max=hi) | |
| out.append(q) | |
| remaining -= block | |
| if remaining > 0: | |
| committed = q[..., q.shape[-1] // 2] | |
| x = torch.cat([x, committed], dim=1) | |
| # Committed medians are OBSERVED context for later blocks. | |
| miss = torch.cat([miss, torch.zeros_like(committed)], dim=1) | |
| return torch.cat(out, dim=1)[:, : int(horizon)] | |
| def forecast_quantiles_batch(self, histories, horizon: int) -> np.ndarray: | |
| x, miss = self._prep(list(histories)) | |
| q = self._decode_quantiles(x, horizon, miss=miss) | |
| qq = q.detach().cpu().numpy().astype(np.float64) | |
| # HURDLE (2026-09-02): zero-movement mixture, horizon-aware. | |
| try: | |
| _H = [np.asarray(h, dtype=np.float64) for h in histories] | |
| for _b, _h in enumerate(_H): | |
| _row = _h[0] if _h.ndim == 2 else _h | |
| qq[_b] = _hurdle_mix(qq[_b], _row) | |
| except Exception: | |
| pass | |
| return qq | |
| def forecast_quantiles(self, history, horizon: int) -> np.ndarray: | |
| return self.forecast_quantiles_batch([history], horizon) | |
| # ── §2: multivariate decode with future-known covariates ───────────────── | |
| # | |
| # The univariate path above cannot express this experiment at all: _prep | |
| # flattens to 1-D and the model call passes no variate_types, so | |
| # roles_on = cfg.use_variate_roles and variate_types is not None | |
| # is False at EVERY inference, and a role-trained checkpoint decodes with | |
| # its role embeddings inert. Everything below exists so a covariate can | |
| # actually reach the model. | |
| def forecast_quantiles_mv(self, histories, horizon: int, *, n_targets: int = 1, | |
| n_future_cov: int = 0, future=None) -> np.ndarray: | |
| """Decode targets given past history and KNOWN future covariates. | |
| histories (B, C, L) channels 0..n_targets-1 are targets; the LAST | |
| n_future_cov are future-known; the rest are | |
| past covariates. Matches assign_roles(), which | |
| is positional by channel index — so the caller | |
| must not permute variates. | |
| future (B, n_future_cov, horizon) the covariate values over the | |
| forecast window. Required when n_future_cov>0: | |
| that knowledge IS the feature being tested. | |
| returns (B, n_targets, horizon, num_q) | |
| Standardisation runs over context AND horizon in one causal pass. That | |
| is deliberate and matches training: causal_standardize excludes masked | |
| entries from its statistics, so the target rows' zero-filled horizon | |
| cannot corrupt their own loc/scale, while the covariate rows — genuinely | |
| observed across the horizon — keep updating exactly as they did during | |
| training. Freezing the covariate scaler at the last context step instead | |
| would introduce a train/inference mismatch that no error would surface. | |
| """ | |
| ps = self.cfg.patch_size | |
| H = int(horizon) | |
| x = np.asarray(histories, dtype=np.float64) | |
| if x.ndim != 3: | |
| raise ValueError(f"histories must be (B, C, L); got {x.shape}") | |
| B, C, L = x.shape | |
| if not 1 <= n_targets <= C: | |
| raise ValueError(f"n_targets={n_targets} outside 1..C={C}") | |
| if n_future_cov < 0 or n_targets + n_future_cov > C: | |
| raise ValueError(f"n_targets={n_targets} + n_future_cov={n_future_cov} > C={C}") | |
| if n_future_cov and future is None: | |
| raise ValueError("n_future_cov>0 requires `future` values") | |
| # Positional roles, mirroring cascade_model.batching.assign_roles. | |
| roles = np.full(C, 1, dtype=np.int64) # ROLE_PAST_COV | |
| roles[:n_targets] = 0 # ROLE_TARGET | |
| if n_future_cov: | |
| roles[C - n_future_cov:] = 2 # ROLE_FUTURE_COV | |
| hz_p = max(1, -(-H // ps)) | |
| n_ctx = max(2, self.cfg.context_length // ps) | |
| if n_ctx + hz_p > self.cfg.max_patches: | |
| n_ctx = self.cfg.max_patches - hz_p | |
| if n_ctx < 2: | |
| raise ValueError( | |
| f"horizon {H} needs {hz_p} patches; max_patches=" | |
| f"{self.cfg.max_patches} leaves no room for context") | |
| window, hz = n_ctx * ps, hz_p * ps | |
| vals = np.zeros((B, C, window + hz), dtype=np.float64) | |
| miss = np.zeros((B, C, window + hz), dtype=np.float64) | |
| for b in range(B): | |
| for c in range(C): | |
| h = x[b, c] | |
| m = ~np.isfinite(h) | |
| if h.shape[0] < window: | |
| obs = h[~m] | |
| fill = obs[0] if obs.size else 0.0 | |
| pad = window - h.shape[0] | |
| h = np.concatenate([np.full(pad, fill), h]) | |
| m = np.concatenate([np.ones(pad, dtype=bool), m]) | |
| else: | |
| h, m = h[-window:], m[-window:] | |
| if m.all(): | |
| h, m = np.zeros_like(h), np.zeros_like(m) | |
| vals[b, c, :window] = np.where(m, 0.0, h) | |
| miss[b, c, :window] = m.astype(np.float64) | |
| miss[:, :, window:] = 1.0 # horizon unobserved … | |
| if n_future_cov: | |
| f = np.asarray(future, dtype=np.float64) | |
| if f.ndim == 2: | |
| f = f[:, None, :] | |
| if f.shape[:2] != (B, n_future_cov): | |
| raise ValueError( | |
| f"future must be (B={B}, n_future_cov={n_future_cov}, >=H); got {f.shape}") | |
| if f.shape[2] < H: | |
| raise ValueError(f"future covers {f.shape[2]} steps; horizon is {H}") | |
| # Pad to the patch boundary by repeating the last known value: those | |
| # steps sit past the requested horizon and are sliced off below. | |
| g = np.concatenate([f[:, :, :H], np.repeat(f[:, :, H - 1:H], hz - H, axis=2)], | |
| axis=2) if hz > H else f[:, :, :hz] | |
| fm = ~np.isfinite(g) | |
| vals[:, C - n_future_cov:, window:] = np.where(fm, 0.0, g) | |
| miss[:, C - n_future_cov:, window:] = fm.astype(np.float64) # … except these | |
| xt = torch.as_tensor(vals, dtype=torch.float64, device=self.device) | |
| mt = torch.as_tensor(miss, dtype=torch.float64, device=self.device) | |
| z, loc, scale = self.m.causal_standardize( | |
| xt.reshape(B * C, -1), mask=mt.reshape(B * C, -1), | |
| binary_passthrough=self.cfg.binary_passthrough) | |
| P = n_ctx + hz_p | |
| patches = z.reshape(B, C, P, ps).to(torch.float32) | |
| pmask = mt.reshape(B, C, P, ps).to(torch.float32) | |
| vt = torch.as_tensor(roles, device=self.device) | |
| pred = self.model(patches, mask=pmask, variate_types=vt) # (B,C,P,ps,Q) | |
| q = pred[:, :n_targets, n_ctx - 1: n_ctx + hz_p - 1] | |
| q, _ = torch.sort(q, dim=-1) | |
| q = q.reshape(B, n_targets, hz, -1)[:, :, :H] | |
| # Invert with the TARGET rows' anchors at the last context step; they are | |
| # frozen across the horizon anyway, since those entries are masked. | |
| loc = loc.reshape(B, C, -1)[:, :n_targets, window - 1][..., None, None] | |
| scale = scale.reshape(B, C, -1)[:, :n_targets, window - 1][..., None, None] | |
| out = torch.sinh(q.double()) * scale.double() + loc.double() | |
| res = out.detach().cpu().numpy().astype(np.float64) | |
| # HURDLE (2026-09-02): zero-movement mixture, horizon-aware. | |
| try: | |
| _H = [np.asarray(h, dtype=np.float64) for h in histories] | |
| for _b, _h in enumerate(_H): | |
| _row = _h[0] if _h.ndim == 2 else _h | |
| res[_b] = _hurdle_mix(res[_b], _row) | |
| except Exception: | |
| pass | |
| return res | |
| def forecast(self, history, horizon: int, num_samples: int) -> np.ndarray: | |
| hist = np.asarray(history, dtype=np.float64).reshape(-1) | |
| seed_src = (hist.tobytes() + int(horizon).to_bytes(8, "big") | |
| + int(num_samples).to_bytes(8, "big")) | |
| seed = int.from_bytes(hashlib.sha256(seed_src).digest()[:8], "big") & ((1 << 63) - 1) | |
| generator = torch.Generator(device=self.device) | |
| generator.manual_seed(seed) | |
| x, miss = self._prep([hist]) | |
| q = self._decode_quantiles(x, horizon, miss=miss)[0] | |
| nq = q.shape[-1] | |
| levels = self.levels | |
| u = torch.rand(int(num_samples), int(horizon), device=self.device, generator=generator) | |
| idx = torch.searchsorted(levels, u.clamp(levels[0].item(), levels[-1].item())) | |
| idx = idx.clamp(1, nq - 1) | |
| qe = q.unsqueeze(0).expand(u.shape[0], -1, -1) | |
| vl = torch.gather(qe, -1, (idx - 1).unsqueeze(-1)).squeeze(-1) | |
| vh = torch.gather(qe, -1, idx.unsqueeze(-1)).squeeze(-1) | |
| ql = levels[idx - 1].double(); qh = levels[idx].double() | |
| frac = ((u.double() - ql) / (qh - ql).clamp_min(1e-8)).clamp(0, 1) | |
| out = vl + frac * (vh - vl) | |
| return out.detach().cpu().numpy().reshape(1, int(num_samples), int(horizon)) | |