yumoto-alpha-18m / forecast_wrapper.py
tensorlink-dev's picture
Upload folder using huggingface_hub
ed7b8dc verified
Raw History Blame Contribute Delete
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
@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))