"""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 @torch.no_grad() 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) @torch.no_grad() 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)] @torch.no_grad() 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. @torch.no_grad() 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 @torch.no_grad() 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))