"""Toto2 backbone + the four upgrades. Forked from ``cascade/trainer/toto2_model.py`` @ cascade main 5e885b3. The fork is deliberate rather than a subclass: the changes thread through ``_Block.forward`` (attention masks), ``Toto2Model.forward`` (variate roles), and the scaler, and this file has to stay **self-contained torch** because it is copied into every checkpoint as ``model.py`` so a wrapper can rebuild the architecture to load weights. No cascade imports. What differs from upstream, and why: * **§0 grouping is honoured per size.** Upstream's ``Toto2Config`` has ``layer_group_size`` but ``from_contract`` never reads it, so every size inherits 4. Every released rung sets it equal to ``num_layers`` (one variate layer each), so a 24-layer model inheriting 4 gets six. * **§2 future-known covariates.** Variate roles are positional — channels ``0..n_targets-1`` are targets, the rest covariates — because the corpus is finite floats with no role axis and the generator contract cannot express observability. Roles drive three things: a learned 3-way type embedding, a per-type time mask (targets and past covariates stay strictly causal, future- known covariates may attend bidirectionally), and an asymmetric variate mask (a covariate query can never read a target key). That trio is what makes target-causality hold by induction over block depth — none of it is specific to a recurrent mixer, which is why this is a mask change and not a backbone rewrite. * **§2b binary detector.** Under the arcsinh scaler a sparse binary column is degenerate: the gap between the two standardised levels blows up as the positive rate goes to zero. Real future-known covariates are mostly binary and sparse (deploy flags, maintenance windows, cron ticks), so binary rows bypass the affine and keep their {0,1} encoding. * **§3 windowed time attention.** A bounded window ``W`` over the patch axis, for the inference-time probe. Set ``time_window = 0`` for unbounded (the default, and what training uses). Everything else — CPM, the robust causal scaler, PerDimScale, xPos, the u-μP residual scheme, the pinball head — is upstream's, unchanged. """ from __future__ import annotations import math import os from dataclasses import dataclass, field import torch import torch.nn as nn import torch.nn.functional as F QUANTILE_LEVELS = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9) # Variate roles. Positional by convention: channels 0..n_targets-1 are targets. ROLE_TARGET = 0 ROLE_PAST_COV = 1 ROLE_FUTURE_COV = 2 N_ROLES = 3 Z_CLAMP = 64.0 @dataclass class CascadeModelConfig: d_model: int = 256 num_layers: int = 4 num_heads: int = 4 head_dim: int = 64 patch_size: int = 32 mlp_expansion: int = 2 d_ff: int = 0 num_quantiles: int = 9 # Toto-2.0-style deep output head: 0 = the original single linear (every # checkpoint before 2026-08-31), >0 = 2-layer MLP head with that hidden # width plus a skip projection. The skip carries the linear-head function, # and linear2 is ZERO-INIT, so at init the MLP head computes exactly what # a fresh linear head would — and a warm-start can drop pretrained linear # head weights into the skip (trainer remaps head.weight -> head.skip.*) # making "add head depth to a trained model" an exact no-op at step 0. head_mlp_hidden: int = 0 # Toto-2.0's exact FFN: bias-free SwiGLU (fc1 emits gate+value, 2*d_ff) # instead of our plain GELU 2-layer. Tensor-verified on the 4m release # (ffn.fc1 [1376, 256] = 2*688, ffn.fc2 [256, 688], no biases). ffn_swiglu: bool = False # Toto-2.0's exact input: skip-projection residual MLP patch->hidden->d # (patch_proj.{linear1,linear2,skip_proj} in the release, hidden 4*d) # replacing our Linear patch_embed + dim-preserving _ResidualMLP. embed_skip_mlp: int = 0 # T4FIX (2026-09-01): MAE-style learned mask token. Fully-missing patches # get this d_model vector INSTEAD of embedding (zeros||mask-flag) — the # nonlinear embed then only ever sees real data. Registered because the # t4 bisect showed the skip-MLP embed learns ~2x worse CPM fill with # overdispersed quantiles when it must embed the missing-patch input # itself. False = off (byte-identical). embed_mask_token: bool = False # Toto-2.0 has NO dim-preserving MLP between the final norm and the head # (only the fused output head) — drop ours for exact replication. no_out_mlp: bool = False context_length: int = 4096 horizon: int = 64 max_patches: int = 256 layer_group_size: int = 4 cpm_c_max: int = 16 cpm_p_max: float = 0.4 residual_mult: float = 0.75 # ── §2 ──────────────────────────────────────────────────────────────────── #: Enable the variate-type embedding and the role-aware masks. Off ⇒ this #: model is numerically upstream's. use_variate_roles: bool = False #: Bypass the arcsinh affine for rows that are binary. Only meaningful with #: roles on, but harmless (and still correct) without them. binary_passthrough: bool = False # ── §3 ──────────────────────────────────────────────────────────────────── #: Bounded time-attention window in PATCHES. 0 = unbounded (training). time_window: int = 0 # Attention sinks: first N patches always visible under a window. # 0 = off, which is bit-identical to the pre-sinks behaviour. attn_sinks: int = 0 #: Attention implementation for the WINDOWED time axis. #: "sdpa" — additive mask into scaled_dot_product_attention. Correct, #: but supplying any attn_mask drops SDPA off its fused kernel and #: onto the materialised-matrix path: measured 0.73x one-shot. #: "flex" — compile the mask INTO a fused kernel via flex_attention. #: The prize is removing that ~27% overhead, NOT block sparsity: at #: ctx 4096 the time axis is 128 patches, where a W=64 window skips only #: 25% of positions and attention is a few percent of a layer anyway. #: Requires torch >= 2.5 AND a compiler toolchain: torch.compile drives #: Triton, which builds a small C extension and therefore needs the #: Python dev headers. A slim image often lacks them, in which case #: compilation raises and this degrades to sdpa (numerically identical, #: verified to ~4e-7 relative) rather than killing the run. To supply #: them without root: #: uv python install 3.12 #: export CPATH=$HOME/.local/share/uv/python/cpython-3.12.*/include/python3.12 #: NOTE the helper links only against libcuda, not libpython, so headers #: from any 3.12.x build are sufficient. attn_impl: str = "sdpa" #: EXP-F control: disable rotary position encoding entirely. The whole #: rope family is downstream of an ARITHMETIC diagnosis (19/32 pairs never #: complete half a turn); this measures whether position encoding is load- #: bearing at all on a 128-position axis. If skill barely moves, the rest of #: the family is a distraction. use_rope: bool = True #: EXP-B: the frequency-ladder base. 10000 was chosen for language contexts #: of many thousands of tokens. rope_scale SHIFTS the ladder uniformly; base #: COMPRESSES it, which is the knob that owns the diagnosed problem — the #: spread from 1.0 to 1.3e-4 rad/position across only 128 positions. #: base ~ 46 makes all 32 pairs complete at least half a turn at L=128, #: versus 13 at stock, and costs nothing at the fast end. rope_base: float = 10000.0 #: Position-interpolation scale on the ROTATION only (see _xpos). 1.0 = stock. #: s < 1 spreads positions across more of the frequency ladder; s > 1 #: compresses them. Not a learned parameter, so it can also be swept at #: inference on a fixed checkpoint. rope_scale: float = 1.0 #: xPos DECAY width. 512 is inherited from a much longer-sequence setting; #: against a 128-position axis the decay exponent only spans +-0.125. This is #: the one part of xPos credited with extrapolation that the rope family #: never swept. 512 reproduces every result recorded so far. xpos_scale_base: float = 512.0 #: YaRN / NTK-by-parts: interpolate SLOW pairs, leave FAST pairs #: extrapolating, instead of PI's uniform rescale. Off reproduces stock. yarn: bool = False #: Ramp bounds in ROTATIONS-per-context. Defaults are set for a 128-position #: axis (r spans ~0.002 to ~20.4), NOT YaRN's published 1/32, which assume #: thousands of tokens and would put every pair below the ramp here — #: silently degrading to plain PI, the method EXP-C measured failing. yarn_alpha: float = 0.5 yarn_beta: float = 8.0 #: YaRN attention temperature. 1.0 = off. attn_temp: float = 1.0 #: PARTIAL ROPE: rotate only the fastest ``k`` frequency pairs and leave the #: rest untouched. 0 = rotate all (stock). #: #: Motivated by EXP-B rather than by the LLM literature. Lowering the base #: (46, 100) made things monotonically WORSE, worst on short horizon. The #: reading: slow pairs act as near-content dims, and compressing the ladder #: forces them to rotate, destroying that. #: #: CAREFUL with the thresholds — an earlier version of this comment conflated #: them and a test caught it. At 128 positions: #: * 13 pairs are USABLE (T*f >= pi, can disambiguate across the context), #: * but only the 7 beyond pair 25 are NEGLIGIBLE (T*f < 0.1). #: Truncating at 13 moves the layer output ~39%; at 25 it moves <2%. So the #: split is a smooth gradient, not a clean 13/19 partition, and the pairs #: between are doing real work. #: #: This makes the allocation explicit and tunable rather than an accident of #: the base. The interesting sweep range is therefore k in ~[13, 32]. rope_partial_k: int = 0 #: LEARNABLE frequency ladder: promote inv_freq from a fixed buffer to a #: parameter (32 scalars per time layer). EXP-B swept ONE degree of freedom #: and found nothing; this gives 32 and lets the model choose its own #: geometry. Also diagnostic — the learned ladder can be read off afterwards, #: which says more than any sweep. Stored in log space so frequencies stay #: positive and the optimiser moves them multiplicatively. rope_learnable: bool = False #: TRAINING-time scale randomisation (§ option 1). 0 = off. Otherwise each #: forward draws rope_scale log-uniformly from [1/j, j], so the model learns #: a scale-INVARIANT distance metric rather than memorising one spacing. #: This is the training-side counterpart to PI, and the direct response to #: EXP-C: you cannot bolt interpolation on at deploy, so train it in. rope_scale_jitter: float = 0.0 #: ARCH 2x2 B-row: the TIME-axis sequence mixer. "attention" (default) is #: the existing xPos MHA and is bit-identical to the pre-field model (no #: extra parameters are created, so old checkpoints load unchanged). #: "mlstm" swaps ONLY the mixing operator inside the same pre-norm / #: depth-scaled-residual scaffold for a TiRex/xLSTM-style matrix-LSTM: #: the stabilized parallel form from NX-AI mlstm_kernels native_stablef #: (logsigmoid forget gates, tril log-decay matrix, row-max stabilizer m, #: qk scale Dh^-0.5, n = max(|sum C~|, exp(-m)) + 1e-6), gate soft-cap 15, #: per-head LayerNorm (eps 1e-6, weight only), sigmoid output gate — #: verified against the published kernel source 2026-08-30. Variate-axis #: blocks stay attention (variates are unordered). Position comes from the #: recurrence, so rope/xPos and window masks do not apply to this mixer; #: future-known-covariate rows get a reversed second pass, averaged, to #: keep the §2 bidirectional contract. time_mixer: str = "attention" @property def ffn_hidden(self) -> int: return self.d_ff if self.d_ff > 0 else self.d_model * self.mlp_expansion def to_dict(self) -> dict: return {k: getattr(self, k) for k in self.__dataclass_fields__} def time_mixer_plan(self) -> list: """Per-layer mixer assignment. None = derive from time_mixer scalar. "pattern:,,..." assigns the k-th TIME block (variate blocks always stay attention) the k-th entry, cycling if the pattern is shorter than the time-block count. H1 (TiRex-2-skeleton hybrid, 2026-08-30): "pattern:mlstm,slstm,mlstm,attention,slstm" — alternating recurrence with one windowed-attention time layer at ~2/3 depth. """ tm = str(getattr(self, "time_mixer", "attention")) if not tm.startswith("pattern:"): return [None] * self.num_layers pat = [x.strip() for x in tm[len("pattern:"):].split(",") if x.strip()] bad = [x for x in pat if x not in ("attention", "mlstm", "slstm", "mamba")] if bad or not pat: raise ValueError(f"bad time_mixer pattern entries: {bad or 'empty'}") plan, k = [], 0 for i in range(self.num_layers): if self.layer_axis(i) == "time": plan.append(pat[k % len(pat)]) k += 1 else: plan.append("attention") return plan def layer_axis(self, i: int) -> str: g = max(1, self.layer_group_size) return "variate" if i % g == g - 1 else "time" @classmethod def from_size(cls, size, **overrides) -> CascadeModelConfig: """Build from a :class:`cascade_model.sizes.Size`, honouring its grouping.""" ctx = int(overrides.pop("context_length", 4096)) hz = int(overrides.pop("horizon", 64)) return cls( d_model=size.d_model, num_layers=size.num_layers, num_heads=size.num_heads, head_dim=size.head_dim, patch_size=size.patch_size, d_ff=size.d_ff, layer_group_size=size.layer_group_size, context_length=ctx, horizon=hz, max_patches=max(8, (ctx + hz) // size.patch_size + 4), **overrides, ) # ── robust causal scaler ───────────────────────────────────────────────────── def is_binary_row(x: torch.Tensor, *, atol: float = 1e-9) -> torch.Tensor: """``(B,)`` bool: which rows of ``(B, L)`` carry only the values 0 and 1. TiRex-2's binary detector. A sparse binary column under an arcsinh scaler is degenerate — with positive rate ``p``, the standardised gap between the two levels grows like ``1/sqrt(p(1-p))`` and diverges as ``p → 0``, so the exact signal a deploy flag carries is the thing the scaler destroys. """ return ((x.abs() < atol) | ((x - 1.0).abs() < atol)).all(dim=-1) def causal_standardize( x: torch.Tensor, mask: torch.Tensor | None = None, *, min_obs: int = 8, eps: float = 1e-5, binary_passthrough: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Per-step causal location/scale under an arcsinh transform. ``x`` is ``(B, L)``; ``mask`` is optional binary ``(B, L)``, 1 = unobserved. Returns ``(z, loc, scale)``, each ``(B, L)``, with ``z = arcsinh((x - loc) / scale)``. With ``binary_passthrough`` a row detected as binary gets ``loc = 0``, ``scale = 1`` — the identity, so ``z = arcsinh(x)`` maps {0,1} to {0, 0.8814} and the encoding survives at any positive rate. """ B, L = x.shape keep = torch.ones_like(x) if mask is None else 1.0 - mask.to(x.dtype) x64 = x.double() k64 = keep.double() ref = x64.gather(-1, (k64 > 0).to(torch.int64).argmax(dim=-1, keepdim=True)) xk = (x64 - ref) * k64 n = k64.cumsum(dim=-1) cnt = n.clamp_min(1.0) loc = xk.cumsum(dim=-1) / cnt var = (xk * xk).cumsum(dim=-1) / cnt - loc * loc loc = loc + ref scale = var.clamp_min(0.0).sqrt().clamp_min(eps) ok = n >= float(min_obs) has = ok.any(dim=-1) first = torch.where( has, ok.to(torch.int64).argmax(dim=-1), torch.full((B,), L - 1, device=x.device) )[:, None] loc = torch.where(ok, loc, loc.gather(-1, first)) scale = torch.where(ok, scale, scale.gather(-1, first)) if binary_passthrough: binr = is_binary_row(x)[:, None] loc = torch.where(binr, torch.zeros_like(loc), loc) scale = torch.where(binr, torch.ones_like(scale), scale) z = torch.asinh((x64 - loc) / scale).clamp_(-Z_CLAMP, Z_CLAMP) return z.to(x.dtype), loc.to(x.dtype), scale.to(x.dtype) def patch_anchors(loc, scale, patch_size: int): B, L = loc.shape P = L // patch_size return ( loc.view(B, P, patch_size)[:, :, -1], scale.view(B, P, patch_size)[:, :, -1], ) def invert_standardize(z, loc, scale): return torch.sinh(z) * scale + loc # ── masks (§2, §3) ─────────────────────────────────────────────────────────── def variate_mask( variate_types: torch.Tensor, group_ids: torch.Tensor | None = None ) -> torch.Tensor: """``(V, V)`` additive mask for the variate-attention axis. Two rules, and the asymmetry is the whole point: * variates only attend within their own group (block-diagonal, as upstream's grouped variate attention already assumes), * a covariate QUERY may never read a target KEY. The second is what preserves target-causality once future-known covariates are allowed to attend bidirectionally along time. Without it, a future-known covariate at horizon position t could read a target key, attend forward in time, and leak a future target value back into an earlier prediction — the exact failure that shows up as a suspiciously good benchmark score. """ v = variate_types.reshape(-1) if group_ids is None: group_ids = torch.zeros_like(v) g = group_ids.reshape(-1) same_group = g.view(-1, 1) == g.view(1, -1) q_is_cov = (v != ROLE_TARGET).view(-1, 1) k_is_tgt = (v == ROLE_TARGET).view(1, -1) allowed = same_group & ~(q_is_cov & k_is_tgt) # A row that can see nothing would produce NaN from softmax over all -inf. # Self-attention is always legal, so pin the diagonal. allowed = allowed | torch.eye(v.numel(), dtype=torch.bool, device=v.device) out = torch.zeros(allowed.shape, dtype=torch.float32, device=v.device) return out.masked_fill(~allowed, float("-inf")) def window_mask(T: int, W: int, *, causal: bool, sinks: int = 0, device=None) -> torch.Tensor: """``(T, T)`` additive mask for a bounded attention window of ``W`` patches. ``causal`` keeps the band strictly at or below the diagonal. ``W <= 0`` means unbounded, in which case the caller should skip the mask entirely and stay on the fast kernel path. ``sinks`` keeps the first ``S`` positions permanently visible to every query IN ADDITION to the sliding band — "attention sinks" (arXiv 2309.17453). Softmax attention concentrates heavily on the earliest positions regardless of their content, so a sliding window that evicts them destabilises the distribution; retaining a handful recovers most of the loss at a fraction of the cache. That matters here because §3 measured a 512-step window failing with a textbook horizon gradient (+3.40 / +6.04 / +9.48% MASE, worst on long), which is the signature this mechanism predicts. ``sinks=0`` is the exact previous behaviour. """ i = torch.arange(T, device=device).view(-1, 1) j = torch.arange(T, device=device).view(1, -1) allowed = (j > i - W) & (j < i + W) if not causal else (j <= i) & (j > i - W) if sinks > 0: sink = j < int(sinks) # A sink is still bound by causality when the axis is causal — a target # row must never read forward, sink or not. allowed = allowed | (sink & (j <= i)) if causal else allowed | sink out = torch.zeros((T, T), dtype=torch.float32, device=device) return out.masked_fill(~allowed, float("-inf")) _FLEX_CACHE: dict = {} _FLEX_FN = None class _FlexUnavailable: """Sentinel: flex was tried and failed; never try again this process.""" def __call__(self, *a, **k): raise RuntimeError("flex unavailable") _FLEX_UNAVAILABLE = _FlexUnavailable() def _additive_from_block(block_mask, q): """Recover a dense additive mask from a BlockMask, for the fallback path.""" dense = block_mask.to_dense() if hasattr(block_mask, "to_dense") else None if dense is None: return None m = dense[0, 0].to(torch.bool) out = torch.zeros(m.shape, dtype=q.dtype, device=q.device) return out.masked_fill(~m, float("-inf")) def _flex_fn(): """``flex_attention``, COMPILED, because uncompiled it defeats the purpose. Called eagerly, flex_attention warns and falls back to an unfused implementation that materialises the full scores matrix — precisely the pessimisation the sdpa+mask path already suffers. Compiling is what turns the mask into a fused kernel and recovers the ~27%. Compiled once and cached at module level: torch.compile has real warmup cost and re-tracing per call would cost far more than the kernel saves. """ global _FLEX_FN if _FLEX_FN is _FLEX_UNAVAILABLE: raise RuntimeError("flex unavailable") if _FLEX_FN is None: from torch.nn.attention.flex_attention import flex_attention _FLEX_FN = torch.compile(flex_attention, dynamic=True) return _FLEX_FN def flex_block_mask(T: int, W: int, *, causal: bool, sinks: int = 0, device=None): """A ``BlockMask`` matching :func:`window_mask`, for ``flex_attention``. Built from the SAME predicate as the additive mask so the two paths cannot drift apart — a fused kernel that quietly attends over a slightly different set than the reference would be indistinguishable from a real result. Cached: constructing a BlockMask is not free and the shape is fixed for a given (T, W, sinks, causal, device). """ from torch.nn.attention.flex_attention import create_block_mask key = (int(T), int(W), int(sinks), bool(causal), str(device)) hit = _FLEX_CACHE.get(key) if hit is not None: return hit Wi, Si = int(W), int(sinks) def mask_mod(b, h, q, kv): band = ((kv > q - Wi) & (kv <= q)) if causal else ((kv > q - Wi) & (kv < q + Wi)) if Si > 0: sink = kv < Si band = band | (sink & (kv <= q)) if causal else band | sink return band bm = create_block_mask(mask_mod, B=None, H=None, Q_LEN=int(T), KV_LEN=int(T), device=device) _FLEX_CACHE[key] = bm return bm # ── building blocks ────────────────────────────────────────────────────────── class _SwiGLU(nn.Module): """Toto-2.0's block FFN: bias-free gated unit, fc1 -> (gate, value).""" def __init__(self, dim: int, hidden: int): super().__init__() self.fc1 = nn.Linear(dim, 2 * hidden, bias=False) self.fc2 = nn.Linear(hidden, dim, bias=False) def forward(self, x): g, v = self.fc1(x).chunk(2, dim=-1) return self.fc2(F.silu(g) * v) class _ResidualMLP(nn.Module): def __init__(self, dim: int, hidden: int): super().__init__() self.net = nn.Sequential( nn.Linear(dim, hidden, bias=False), nn.SiLU(), nn.Linear(hidden, dim, bias=False), ) def forward(self, x): return x + self.net(x) class _MLPHead(nn.Module): """Toto-2.0-style deep output head: skip(x) + linear2(act(linear1(x))). Datadog's checkpoint puts 1.79M params here (512 -> 2048 -> 288 + skip) where our recipe has always used the 0.15M single linear — the largest structural difference between the two models, targeted at distribution shape, which is where ALL measured fine-tune value lives (MASE flat, CRPS gains). linear2 is zero-init so the head starts as exactly the skip (i.e. exactly a linear head): from scratch that is the standard init, and a warm-start that maps pretrained head weights onto the skip is an exact function-preserving upgrade. """ def __init__(self, dim: int, hidden: int, out: int): super().__init__() self.skip = nn.Linear(dim, out) self.linear1 = nn.Linear(dim, hidden) self.act = nn.SiLU() self.linear2 = nn.Linear(hidden, out) def forward(self, x): return self.skip(x) + self.linear2(self.act(self.linear1(x))) def _yarn_inv_freq(inv_freq, rope_scale: float, n_pos: int, alpha: float, beta: float): """YaRN / NTK-by-parts: interpolate SLOW pairs, leave FAST pairs alone. Plain PI divides every position by ``s``, so every relative distance in the sequence is rescaled at once. Measured (EXP-C), that is much worse zero-shot than not extending at all: -12.8 / -25.3 / -24.1% against naive at s=2. YaRN's argument is that the two ends of the ladder want opposite treatment. A pair that completes many rotations across the context is carrying local, high-resolution distance information and should be left EXTRAPOLATING. A pair that has barely turned is carrying absolute-ish position and is the one that actually goes out of range, so it should be INTERPOLATED. The ramp runs on ``r_i``, the number of full rotations pair ``i`` completes across ``n_pos`` positions:: r_i = n_pos * inv_freq_i / (2 * pi) r_i <= alpha -> fully interpolated (inv_freq / s) r_i >= beta -> untouched (inv_freq) between -> linear blend **alpha/beta must be set for THIS axis, not copied from the LLM defaults.** YaRN's published 1/32 assume thousands of tokens. Here the whole axis is 128 positions and r spans ~0.002 to ~20.4, so with beta=32: * 21 of 32 pairs sit at or below alpha and are FULLY interpolated, * the remaining 11 are only partially ramped, * and NO pair ever reaches full extrapolation — the fastest tops out at gamma ~ 0.63. That is not identical to plain PI, but it leans heavily toward it, and plain PI is the method EXP-C measured failing. The defaults below put the ramp where this ladder actually lives. """ r = n_pos * inv_freq / (2.0 * math.pi) if beta <= alpha: raise ValueError(f"yarn beta({beta}) must exceed alpha({alpha})") gamma = ((r - alpha) / (beta - alpha)).clamp(0.0, 1.0) # 1 = extrapolate return gamma * inv_freq + (1.0 - gamma) * (inv_freq / rope_scale) def _xpos(q, k, inv_freq, zeta, scale_base: float = 512.0, rope_scale: float = 1.0, *, yarn: bool = False, yarn_alpha: float = 0.5, yarn_beta: float = 8.0, attn_temp: float = 1.0, partial_k: int = 0): """xPos = RoPE rotation + a per-dimension decay. ``rope_scale`` divides the position before the rotation (position interpolation). Deployed, PI uses s > 1 to compress an over-long axis back into the trained range. TRAINED, the interesting direction is s < 1, which SPREADS positions across more of the frequency ladder. The motivation is that the ladder is badly matched to this axis. With head_dim=64, base=10000 and a 4096-step context at patch_size=32, the time axis is only 128 positions: the fastest pair sweeps 20 cycles while the slowest sweeps 1.0 degree end-to-end, and **19 of 32 pairs never complete half a rotation**, so they carry no within-context positional information. Under the W=2048 production window (64 positions) it is 21 of 32. s = 1/2 doubles every rate, activating 3 more pairs while keeping the fastest at 2 rad/position — clear of the aliasing wall at s < 1/pi ~ 0.318. The DECAY term is a SEPARATE mechanism from the rotation — attention falloff with distance — and ``scale_base`` is its width. It is also the part of xPos that the original paper credits for extrapolation, and this project never tuned it: the whole rope family (EXP-B base, EXP-C scale) swept rotation and left decay at its default. ``scale_base=512`` is inherited from a setting with sequences several times longer than ours. Against 128 positions the exponent ``(t - T//2)/512`` only spans +-0.125, so the decay operates in a heavily compressed corner of its range. Lowering it widens that range. NOTE this is arithmetic plus reasoning, NOT a published result — I could find no literature tuning scale_base as a function of sequence length, so it is a hypothesis with a cheap test. """ T = q.shape[-2] t = torch.arange(T, device=q.device, dtype=inv_freq.dtype) if yarn: eff = _yarn_inv_freq(inv_freq, rope_scale, T, yarn_alpha, yarn_beta) freqs = torch.outer(t, eff) else: freqs = torch.outer(t / rope_scale, inv_freq) if partial_k: # Rotate only the fastest `partial_k` pairs; zero the rest so cos=1, # sin=0 and those dims pass through unrotated. Cheaper and clearer than # slicing the tensors, and it keeps every downstream shape identical. if not 0 < partial_k <= freqs.shape[-1]: raise ValueError( f"rope_partial_k={partial_k} outside 1..{freqs.shape[-1]}") freqs = freqs.clone() freqs[:, partial_k:] = 0.0 cos = freqs.cos().repeat_interleave(2, dim=-1) sin = freqs.sin().repeat_interleave(2, dim=-1) power = ((t - T // 2) / scale_base)[:, None] scale = (zeta[None, :] ** power).repeat_interleave(2, dim=-1) def rotate(x): x1 = x[..., 0::2] x2 = x[..., 1::2] return torch.stack((-x2, x1), dim=-1).flatten(-2) qo = (q * cos + rotate(q) * sin) * scale ko = (k * cos + rotate(k) * sin) / scale if attn_temp != 1.0: # YaRN's second half: a longer context spreads softmax mass thinner, so # it sharpens the logits by a constant. Folded into q because the logits # are q·k — equivalent, and it keeps the fused attention kernel. qo = qo * attn_temp return qo, ko def _soft_cap(x: torch.Tensor, cap: float = 15.0) -> torch.Tensor: """xLSTM gate soft-cap: cap * tanh(x / cap).""" return cap * torch.tanh(x / cap) class _MultiHeadNorm(nn.Module): """Per-head LayerNorm, weight only (TiRex MultiHeadLayerNorm, eps 1e-6).""" def __init__(self, num_heads: int, head_dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(num_heads, head_dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: # (N, H, T, Dh) mu = x.mean(dim=-1, keepdim=True) var = x.var(dim=-1, keepdim=True, unbiased=False) return (x - mu) / torch.sqrt(var + self.eps) * self.weight[None, :, None, :] def _mlstm_scan(q, k, v, i_pre, f_pre, eps: float = 1e-6): """Stabilized parallel mLSTM — the exact native_stablef math. ``q, k, v``: ``(N, H, T, Dh)``; ``i_pre, f_pre``: ``(N, H, T)`` soft-capped gate preactivations. The decay matrix is built as the cumsum DIFFERENCE ``Fc[t] - Fc[s] + i[s]`` rather than the kernel's repeat/tril/cumsum — identical values (diagonal reduces to ``i[t]``), and with the +-15 soft cap the cancellation is bounded by ~15*T, well inside float32. At T = 128 the quadratic form costs a few MB and needs no chunking. """ N, H, T, Dh = q.shape logf = F.logsigmoid(f_pre) # (N, H, T) fc = logf.cumsum(dim=-1) D = fc[..., :, None] - fc[..., None, :] + i_pre[..., None, :] tril = torch.ones(T, T, dtype=torch.bool, device=q.device).tril() D = D.masked_fill(~tril, float("-inf")) m = D.max(dim=-1, keepdim=True).values # (N, H, T, 1) Dm = torch.exp(D - m) # diag finite => m finite S = (q @ k.transpose(-2, -1)) * (Dh ** -0.5) Ct = S * Dm n = torch.maximum(Ct.sum(dim=-1, keepdim=True).abs(), torch.exp(-m)) return (Ct / (n + eps)) @ v def _mamba_scan(q, k, v, dt, A_log, skip): """Mamba-2 SSD, dual (decay-masked attention) form — Dao & Gu 2024, eq. 5. ``q``: C (readout), ``k``: B (write key), ``v``: x (values), all ``(N, H, T, Dh)`` with d_state = head_dim; ``dt``: ``(N, H, T)`` softplus'd step sizes; ``A_log``: ``(H,)`` log of the positive decay rate; ``skip``: ``(H,)`` the D residual. The recurrence h_t = exp(-dt_t*A) h_{t-1} + dt_t B_t x_t^T, y_t = C_t h_t + D x_t collapses at our T (~130) to one masked quadratic form, exactly the shape of _mlstm_scan's — with two simplifications the math hands us: log-decay is <= 0 everywhere so exp(D) <= 1 and no row-max stabilizer is needed, and there is no normalizer n (SSD is not softmax-normalized; magnitude lives in B/C/dt). The official Triton kernels only pay at multi-thousand-token sequences; at 2 chunks of 64 this matmul form IS the efficient implementation. """ N, H, T, Dh = q.shape loga = -A_log.exp()[None, :, None] * dt # (N,H,T) <= 0 fc = loga.cumsum(dim=-1) D = fc[..., :, None] - fc[..., None, :] # decay j+1..i tril = torch.ones(T, T, dtype=torch.bool, device=q.device).tril() Dm = torch.exp(D.masked_fill(~tril, float("-inf"))) S = (q @ k.transpose(-2, -1)) * Dm xbar = v * dt[..., None] # Δ-discretized return S @ xbar + skip[None, :, None, None] * v def _slstm_scan_impl(x_gates, R, h0=None, eps: float = 1e-6): """Stabilized sLSTM (TiRex's mixer — the state-tracking xLSTM cell). ``x_gates``: ``(4, N, T, H, Dh)`` input-side gate preactivations in order (i, f, z, o); ``R``: ``(4, H, Dh, Dh)`` per-head block-diagonal recurrent weights applied to h_{t-1} — the NON-DIAGONAL recurrence that no parallel form can express, which is the whole point of this cell. Sequential over T by necessity; at T = 128 that is 128 small batched einsums. Paper math (Beck et al. 2024), sigmoid-forget variant in log space: m_t = max(logsigmoid(f~) + m_{t-1}, i~) i' = exp(i~ - m_t); f' = exp(logsigmoid(f~) + m_{t-1} - m_t) c_t = f' c_{t-1} + i' tanh(z~); n_t = f' n_{t-1} + i' h_t = sigmoid(o~) * c_t / (n_t + eps) """ _, N, T, H, Dh = x_gates.shape dev, dt = x_gates.device, x_gates.dtype c = torch.zeros(N, H, Dh, device=dev, dtype=dt) n = torch.zeros(N, H, Dh, device=dev, dtype=dt) m = torch.full((N, H, Dh), -1e9, device=dev, dtype=dt) h = torch.zeros(N, H, Dh, device=dev, dtype=dt) if h0 is None else h0 out = torch.empty(N, T, H, Dh, device=dev, dtype=dt) for t in range(T): rec = torch.einsum("nhd,ghde->gnhe", h, R) # (4, N, H, Dh) i_pre = x_gates[0, :, t] + rec[0] f_pre = x_gates[1, :, t] + rec[1] z = torch.tanh(x_gates[2, :, t] + rec[2]) o = torch.sigmoid(x_gates[3, :, t] + rec[3]) logf = F.logsigmoid(f_pre) m_new = torch.maximum(logf + m, i_pre) i_s = torch.exp(i_pre - m_new) f_s = torch.exp(logf + m - m_new) c = f_s * c + i_s * z n = f_s * n + i_s m = m_new h = o * c / (n + eps) out[:, t] = h return out.permute(0, 2, 1, 3) # (N, H, T, Dh) # The naive loop is KERNEL-LAUNCH bound, not FLOP bound: 128 steps x ~12 tiny # ops per time block measured 132K tok/s on an H-class card (13x under the # mLSTM cell). torch.compile fuses the pointwise chain and shrinks the launch # count several-fold; compiled lazily and per-process, with a hard fallback to # the eager loop when the box lacks Triton/dev headers (same degradation # policy as flex_attention above). dynamic=False: the mixed channel mode # yields a small fixed set of batch shapes, each compiled once. _SLSTM_SCAN_FN = None def _slstm_scan(x_gates, R, h0=None, eps: float = 1e-6): global _SLSTM_SCAN_FN if _SLSTM_SCAN_FN is None: try: # dynamic=False compiles once per shape — right for training (a # small fixed shape set) but a recompile STORM for the GIFT eval, # which sweeps ~97 config shapes (measured: 2h+ eval instead of # ~25 min, all of it gcc). Eval harnesses set # CASCADE_SLSTM_DYNAMIC=1 to compile shape-polymorphic instead. _dyn = os.environ.get("CASCADE_SLSTM_DYNAMIC", "") == "1" _SLSTM_SCAN_FN = torch.compile(_slstm_scan_impl, dynamic=_dyn) except Exception: _SLSTM_SCAN_FN = _slstm_scan_impl if _SLSTM_SCAN_FN is not _slstm_scan_impl: try: return _SLSTM_SCAN_FN(x_gates, R, h0, eps) except Exception: _SLSTM_SCAN_FN = _slstm_scan_impl return _slstm_scan_impl(x_gates, R, h0, eps) class _Block(nn.Module): """Pre-norm MHA + GELU MLP, over either the time or the variate axis. ``axis="time"``: causal over the patch axis with xPos positions — except for rows flagged bidirectional (future-known covariates), and except for a bounded window when ``time_window > 0``. ``axis="variate"``: full attention over the variate axis (no positions — variates are unordered), optionally under an asymmetric role mask. """ def __init__(self, cfg: CascadeModelConfig, axis: str, block_idx: int = 0, mixer: str | None = None): super().__init__() self.cfg = cfg self.axis = axis inner = cfg.num_heads * cfg.head_dim if mixer is None: req = str(getattr(cfg, "time_mixer", "attention")) mixer = req if req in ("mlstm", "slstm", "mamba") else "attention" self.mixer = mixer if axis == "time" else "attention" self.norm1 = nn.LayerNorm(cfg.d_model, eps=1e-4, elementwise_affine=False) if self.mixer != "slstm": self.qkv = nn.Linear(cfg.d_model, 3 * inner, bias=True) self.proj = nn.Linear(inner, cfg.d_model, bias=True) self.norm2 = nn.LayerNorm(cfg.d_model, eps=1e-4, elementwise_affine=False) hidden = cfg.ffn_hidden if cfg.ffn_swiglu: self.mlp = _SwiGLU(cfg.d_model, hidden) else: self.mlp = nn.Sequential( nn.Linear(cfg.d_model, hidden, bias=False), nn.GELU(), nn.Linear(hidden, cfg.d_model, bias=False), ) if self.mixer == "slstm": # TiRex's mixer: 4 gates (i,f,z,o), per-dim, input side from the # normed token + per-head block-diagonal recurrence on h_{t-1}. # No qkv — the cell IS the mixing. Inits set in reset_gate_biases. self.gates_x = nn.Linear(cfg.d_model, 4 * inner, bias=True) self.rec = nn.Parameter(torch.zeros(4, cfg.num_heads, cfg.head_dim, cfg.head_dim)) self.mh_norm = _MultiHeadNorm(cfg.num_heads, cfg.head_dim) elif self.mixer == "mlstm": # Gate heads follow xLSTM-large: i and f are per-head scalars from # the normed input; o is elementwise over the inner dim. per_dim_scale # and the rope buffers are attention-specific and deliberately NOT # created here, so attention-mode state_dicts stay byte-identical. self.gates_if = nn.Linear(cfg.d_model, 2 * cfg.num_heads, bias=True) self.ogate = nn.Linear(cfg.d_model, inner, bias=True) self.mh_norm = _MultiHeadNorm(cfg.num_heads, cfg.head_dim) elif self.mixer == "mamba": # Mamba-2 head: qkv doubles as (C, B, x); per-head step size dt, # per-head decay rate A (stored in log), per-head D skip, SiLU # z-gate on the inner dim (reuses the ogate module name so the # trainer's no-decay match covers its bias too). Inits land in # reset_gate_biases. self.mamba_dt = nn.Linear(cfg.d_model, cfg.num_heads, bias=True) self.mamba_A_log = nn.Parameter(torch.zeros(cfg.num_heads)) self.mamba_skip = nn.Parameter(torch.ones(cfg.num_heads)) self.ogate = nn.Linear(cfg.d_model, inner, bias=True) self.mh_norm = _MultiHeadNorm(cfg.num_heads, cfg.head_dim) else: self.per_dim_scale = nn.Parameter(torch.zeros(cfg.head_dim)) if axis == "time" and self.mixer == "attention": half = cfg.head_dim // 2 idx = torch.arange(half).float() / max(1, half) base_freq = 1.0 / (cfg.rope_base**idx) if cfg.rope_learnable: # log space: keeps frequencies positive under any update and # makes gradient steps multiplicative, which is the right metric # for a ladder spanning four orders of magnitude. Initialised at # the stock ladder, so step 0 is exactly stock. self.log_inv_freq = nn.Parameter(base_freq.log()) self.register_buffer("inv_freq", base_freq, persistent=False) else: self.register_buffer("inv_freq", base_freq, persistent=False) self.register_buffer("zeta", (idx + 0.4) / 1.4, persistent=False) S = max(2.0, cfg.context_length / cfg.patch_size) ratio2 = S / math.log(S) af2 = 2.0 * cfg.residual_mult**2 / (ratio2 + 1.0) aa2 = ratio2 * af2 L = 2.0 * cfg.num_layers i = block_idx tau2_attn = aa2 / (L / 2.0 + i * aa2 + i * af2) tau2_mlp = af2 / (L / 2.0 + (i + 1) * aa2 + i * af2) self.attn_a = math.sqrt(tau2_attn / (tau2_attn + 1.0)) self.attn_b = math.sqrt(1.0 / (tau2_attn + 1.0)) self.mlp_a = math.sqrt(tau2_mlp / (tau2_mlp + 1.0)) self.mlp_b = math.sqrt(1.0 / (tau2_mlp + 1.0)) def reset_gate_biases(self) -> None: """xLSTM-7B gate init, applied AFTER the model-wide bias zeroing. Forget bias linspace(3, 6) per head starts the memory near-preserving (logsigmoid(3..6) ~ -0.05..-0.002); input bias -10 starts writes near-off. Both are inside the +-15 soft cap. Without this, exp input gates at bias 0 make every position write at full strength from step 0. """ H, Dh = self.cfg.num_heads, self.cfg.head_dim inner = H * Dh if self.mixer == "mlstm": with torch.no_grad(): self.gates_if.bias[:H].fill_(-10.0) self.gates_if.bias[H:].copy_(torch.linspace(3.0, 6.0, H)) elif self.mixer == "slstm": # xlstm "small_init" convention: forget bias linspace(3,6) per dim, # input bias -10 (writes start near-off), z/o biases zero, # recurrent kernel zeros (their default) — the cell starts as a # feedforward gate and learns its recurrence. with torch.no_grad(): self.gates_x.bias[:inner].fill_(-10.0) self.gates_x.bias[inner:2 * inner].copy_( torch.linspace(3.0, 6.0, Dh).repeat(H)) self.gates_x.bias[2 * inner:].zero_() elif self.mixer == "mamba": # Mamba-2 inits, deterministic variants of the paper's draws: # A = linspace(1,16) per head (their U[1,16]); dt_bias = inverse # softplus of a log-spaced dt in [1e-3, 1e-1] so heads start with # a spread of timescales spanning ~10 to ~1000 steps of memory. with torch.no_grad(): A0 = torch.linspace(1.0, 16.0, H) self.mamba_A_log.copy_(A0.log()) dt0 = torch.logspace(math.log10(1e-3), math.log10(1e-1), H) self.mamba_dt.bias.copy_(torch.log(torch.expm1(dt0))) def _mix_slstm(self, x, h, *, bidirectional_rows=None): """TiRex-style sLSTM time mixing inside the host residual scaffold. Same contract as _mix_mlstm: rope/window masks do not apply (position and memory live in the recurrence); future-known covariate rows get a time-reversed second pass, averaged. """ N, T, _ = x.shape H, Dh = self.cfg.num_heads, self.cfg.head_dim g = self.gates_x(h).view(N, T, 4, H, Dh).permute(2, 0, 1, 3, 4) out = _slstm_scan(g, self.rec) # (N, H, T, Dh) if bidirectional_rows is not None and bool(bidirectional_rows.any()): bid = bidirectional_rows rev = _slstm_scan(g[:, bid].flip(2), self.rec).flip(2) out = out.clone() out[bid] = 0.5 * (out[bid] + rev) mixed = self.mh_norm(out).transpose(1, 2).reshape(N, T, H * Dh) x = self.attn_b * x + self.attn_a * self.proj(mixed) return self.mlp_b * x + self.mlp_a * self.mlp(self.norm2(x)) def _mix_mamba(self, x, h, *, bidirectional_rows=None): """Mamba-2 time mixing inside the host residual scaffold. Same contract as the other recurrent mixers: rope/window masks do not apply (position lives in the decay), and future-known covariate rows get a time-reversed second pass, averaged. """ N, T, _ = x.shape H, Dh = self.cfg.num_heads, self.cfg.head_dim qkv = self.qkv(h).view(N, T, 3, H, Dh) q, k, v = (t.transpose(1, 2) for t in qkv.unbind(dim=2)) # (N,H,T,Dh) dt = F.softplus(self.mamba_dt(h)).transpose(1, 2) # (N,H,T) out = _mamba_scan(q, k, v, dt, self.mamba_A_log, self.mamba_skip) if bidirectional_rows is not None and bool(bidirectional_rows.any()): bid = bidirectional_rows rev = _mamba_scan( q[bid].flip(2), k[bid].flip(2), v[bid].flip(2), dt[bid].flip(2), self.mamba_A_log, self.mamba_skip, ).flip(2) out = out.clone() out[bid] = 0.5 * (out[bid] + rev) mixed = self.mh_norm(out).transpose(1, 2).reshape(N, T, H * Dh) mixed = mixed * F.silu(self.ogate(h)) x = self.attn_b * x + self.attn_a * self.proj(mixed) return self.mlp_b * x + self.mlp_a * self.mlp(self.norm2(x)) def _mix_mlstm(self, x, h, *, bidirectional_rows=None): """TiRex-style mLSTM time mixing inside the host residual scaffold. ``h`` is the pre-normed input. Ignores rope and window masks (position and locality live in the recurrence); bidirectional rows (future-known covariates) get a time-reversed second pass, averaged. """ N, T, _ = x.shape H, Dh = self.cfg.num_heads, self.cfg.head_dim qkv = self.qkv(h).view(N, T, 3, H, Dh) q, k, v = (t.transpose(1, 2) for t in qkv.unbind(dim=2)) # (N,H,T,Dh) g = _soft_cap(self.gates_if(h)) # (N,T,2H) i_pre = g[..., :H].transpose(1, 2) # (N,H,T) f_pre = g[..., H:].transpose(1, 2) out = _mlstm_scan(q, k, v, i_pre, f_pre) if bidirectional_rows is not None and bool(bidirectional_rows.any()): bid = bidirectional_rows rev = _mlstm_scan( q[bid].flip(2), k[bid].flip(2), v[bid].flip(2), i_pre[bid].flip(2), f_pre[bid].flip(2), ).flip(2) out = out.clone() out[bid] = 0.5 * (out[bid] + rev) mixed = self.mh_norm(out).transpose(1, 2).reshape(N, T, H * Dh) mixed = mixed * torch.sigmoid(self.ogate(h)) x = self.attn_b * x + self.attn_a * self.proj(mixed) return self.mlp_b * x + self.mlp_a * self.mlp(self.norm2(x)) def _attend(self, q, k, v, *, causal: bool, attn_mask=None, block_mask=None): if block_mask is not None: # scale must match the SDPA path exactly — this model uses 1/d, not # the conventional 1/sqrt(d), and a mismatch here would look like a # subtle quality regression rather than a bug. try: return _flex_fn()(q, k, v, block_mask=block_mask, scale=1.0 / self.cfg.head_dim) except Exception: # Compilation can fail for reasons that have nothing to do with # this model — Triton builds a small C extension and needs the # Python dev headers, which a slim image may not carry. The two # paths are numerically identical (verified to ~3e-7 relative), # so degrade to sdpa rather than kill a training run hours in. global _FLEX_FN _FLEX_FN = _FLEX_UNAVAILABLE attn_mask = _additive_from_block(block_mask, q) return F.scaled_dot_product_attention( q, k, v, is_causal=(causal and attn_mask is None), attn_mask=attn_mask, scale=1.0 / self.cfg.head_dim, ) def forward(self, x, *, bidirectional_rows=None, attn_mask=None, attn_mask_bidir=None, block_mask=None, block_mask_bidir=None): """``x`` is ``(N, T, d)``. ``bidirectional_rows`` (time axis only) is an ``(N,)`` bool selecting rows that may attend forward — future-known covariates. Rather than materialise a ``(V, T, T)`` mask, the batch is SPLIT by type and the two halves run as separate calls, which keeps both on the fast kernel path. ``attn_mask`` is an additive ``(T, T)`` applied to every row. """ N, T, _ = x.shape h = self.norm1(x) if self.mixer == "slstm": return self._mix_slstm(x, h, bidirectional_rows=bidirectional_rows) if self.mixer == "mlstm": return self._mix_mlstm(x, h, bidirectional_rows=bidirectional_rows) if self.mixer == "mamba": return self._mix_mamba(x, h, bidirectional_rows=bidirectional_rows) qkv = self.qkv(h).view(N, T, 3, self.cfg.num_heads, self.cfg.head_dim) q, k, v = qkv.unbind(dim=2) q, k, v = (t.transpose(1, 2) for t in (q, k, v)) if self.axis == "time" and self.cfg.use_rope: if self.cfg.rope_learnable: # Nyquist guard. A rotation above pi radians PER POSITION # aliases: adjacent positions become indistinguishable and the # pair emits noise rather than position. Nothing in the loss # prevents the optimiser walking there, and the failure is # silent — the model would just get worse for a reason no # metric names. Clamped in log space, where the parameter lives. inv_freq = self.log_inv_freq.clamp(max=math.log(math.pi)).exp() else: inv_freq = self.inv_freq s = self.cfg.rope_scale j = self.cfg.rope_scale_jitter if j and self.training: # Log-uniform in [1/j, j] so shrink and stretch are symmetric — # uniform in s would bias every draw toward compression. One draw # per forward, not per row: within a batch the positional metric # must be consistent or attention compares incompatible spacings. u = torch.rand((), device=q.device).item() s = s * float(math.exp((2.0 * u - 1.0) * math.log(j))) q, k = _xpos(q, k, inv_freq, self.zeta, scale_base=self.cfg.xpos_scale_base, rope_scale=s, yarn=self.cfg.yarn, yarn_alpha=self.cfg.yarn_alpha, yarn_beta=self.cfg.yarn_beta, attn_temp=self.cfg.attn_temp, partial_k=self.cfg.rope_partial_k) q = q * (F.softplus(self.per_dim_scale) / math.log(2.0)) causal = self.axis == "time" # Bidirectional rows need their OWN mask. `_attend` sets # is_causal=(causal and attn_mask is None), so once a window mask is # supplied, causality comes from the MASK alone — and window_mask() is # built causal. Feeding the causal band to the bidirectional half # therefore makes future-known covariates causal, silently: no error, no # metric, just the §2 capability quietly gone. Verified by probe — a # future-cov row's dependence on a future position drops from 3.8e-2 to # exactly 0 the moment a window is enabled. bmask = attn_mask_bidir if attn_mask_bidir is not None else ( None if attn_mask is None else attn_mask ) # A block mask supersedes the additive one: it encodes the SAME # predicate, causality included, so passing both would be redundant and # passing the additive one alongside would force the slow path anyway. bb = block_mask_bidir if block_mask_bidir is not None else block_mask am = None if block_mask is not None else attn_mask ab = None if bb is not None else bmask if bidirectional_rows is None or not causal: attn = self._attend(q, k, v, causal=causal, attn_mask=am, block_mask=block_mask) elif bool(bidirectional_rows.all()): attn = self._attend(q, k, v, causal=False, attn_mask=ab, block_mask=bb) elif not bool(bidirectional_rows.any()): attn = self._attend(q, k, v, causal=True, attn_mask=am, block_mask=block_mask) else: bid = bidirectional_rows attn = torch.empty_like(q) attn[~bid] = self._attend( q[~bid], k[~bid], v[~bid], causal=True, attn_mask=am, block_mask=block_mask, ) # Future-known covariates: bidirectional along time. Safe only # because variate_mask() forbids their queries from reading target # keys — otherwise this is the leak. attn[bid] = self._attend( q[bid], k[bid], v[bid], causal=False, attn_mask=ab, block_mask=bb, ) attn = attn.transpose(1, 2).reshape(N, T, self.cfg.num_heads * self.cfg.head_dim) x = self.attn_b * x + self.attn_a * self.proj(attn) x = self.mlp_b * x + self.mlp_a * self.mlp(self.norm2(x)) return x class CascadeModel(nn.Module): """Patch transformer with CPM, alternating time/variate attention, and the §2/§3 role machinery. Predicts each position's NEXT patch as quantiles.""" def __init__(self, cfg: CascadeModelConfig): super().__init__() self.cfg = cfg if cfg.embed_skip_mlp > 0: # Toto-2.0 patch_proj: skip(64->d) + linear2(act(linear1(64->h))) self.patch_embed = _MLPHead(cfg.patch_size * 2, cfg.embed_skip_mlp, cfg.d_model) self.embed_mlp = nn.Identity() else: self.patch_embed = nn.Linear(cfg.patch_size * 2, cfg.d_model) self.embed_mlp = _ResidualMLP(cfg.d_model, cfg.ffn_hidden) # §2: learned 3-way role embedding, added AFTER the residual-MLP patch # projection so it colours the token the transformer sees, not the raw # patch. Zero-init ⇒ enabling roles starts as an exact no-op. self.role_embed = nn.Embedding(N_ROLES, cfg.d_model) if cfg.use_variate_roles else None # Zero-init: at step 0 a masked patch's token is the origin, close to # what the linear embed's bias-only output would be — trains freely. self.mask_token = (nn.Parameter(torch.zeros(cfg.d_model)) if cfg.embed_mask_token else None) _plan = cfg.time_mixer_plan() self.blocks = nn.ModuleList( _Block(cfg, axis=cfg.layer_axis(i), block_idx=i, mixer=_plan[i]) for i in range(cfg.num_layers) ) self.norm = nn.LayerNorm(cfg.d_model, eps=1e-4, elementwise_affine=False) self.out_mlp = (nn.Identity() if cfg.no_out_mlp else _ResidualMLP(cfg.d_model, cfg.ffn_hidden)) if cfg.head_mlp_hidden > 0: self.head = _MLPHead(cfg.d_model, cfg.head_mlp_hidden, cfg.patch_size * cfg.num_quantiles) else: self.head = nn.Linear(cfg.d_model, cfg.patch_size * cfg.num_quantiles) self.apply(self._init_weights) if self.role_embed is not None: nn.init.zeros_(self.role_embed.weight) if cfg.head_mlp_hidden > 0: # AFTER the global init: the zero-init contract in _MLPHead's # docstring only holds if nothing re-randomises linear2. nn.init.zeros_(self.head.linear2.weight) nn.init.zeros_(self.head.linear2.bias) for blk in self.blocks: blk.reset_gate_biases() # no-op on attention blocks def _init_weights(self, m: nn.Module) -> None: if isinstance(m, nn.Linear): fan_in = m.weight.shape[1] nn.init.normal_(m.weight, mean=0.0, std=1.0 / math.sqrt(fan_in)) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.Embedding): nn.init.normal_(m.weight, mean=0.0, std=0.02) def forward( self, patches: torch.Tensor, mask: torch.Tensor | None = None, *, variate_types: torch.Tensor | None = None, group_ids: torch.Tensor | None = None, ) -> torch.Tensor: """``patches``: ``(B, P, ps)`` or ``(B, C, P, ps)``. ``mask``: binary, 1 = unobserved, patch-level or per-entry. ``variate_types``: ``(C,)`` in {0 target, 1 past-cov, 2 future-known-cov}. Returns ``(B, [C,] P, ps, num_q)``.""" squeeze_variates = patches.dim() == 3 if squeeze_variates: patches = patches[:, None] if mask is not None: mask = mask[:, None] B, C, P, ps = patches.shape if mask is None: mask = torch.zeros_like(patches) else: if mask.dim() == 3: mask = mask[..., None].expand(B, C, P, ps) mask = mask.to(patches.dtype) x = torch.cat([patches * (1.0 - mask), mask], dim=-1) x = self.embed_mlp(self.patch_embed(x)) # (B, C, P, d) if self.mask_token is not None: # Patch-level replacement only when EVERY entry is missing — # partially observed patches keep the embed path (it still sees # real values there). full = mask.mean(dim=-1, keepdim=True) >= 1.0 - 1e-6 x = torch.where(full, self.mask_token.to(x.dtype).view(1, 1, 1, -1), x) roles_on = self.cfg.use_variate_roles and variate_types is not None if roles_on: vt = variate_types.to(x.device).long().reshape(-1) if vt.numel() != C: raise ValueError(f"variate_types has {vt.numel()} entries; expected C={C}") x = x + self.role_embed(vt).view(1, C, 1, -1) # Row b*C+c has the type of channel c — matches the reshape below. bidir_rows = (vt == ROLE_FUTURE_COV).repeat(B) vmask = variate_mask(vt, group_ids) else: bidir_rows = None vmask = None W = int(self.cfg.time_window) S = int(getattr(self.cfg, "attn_sinks", 0)) tmask = (window_mask(P, W, causal=True, sinks=S, device=x.device) if W > 0 else None) # The SYMMETRIC band, for future-known covariate rows only: they are # bounded by the same window but may look forward within it. Built only # when such rows exist, so the univariate path allocates nothing. tmask_bidir = ( window_mask(P, W, causal=False, sinks=S, device=x.device) if W > 0 and bidir_rows is not None and bool(bidir_rows.any()) else None ) # flex_attention only earns its keep when a mask is needed at all: with # W=0 plain SDPA already takes the fused causal path and is the fastest # option, so this never engages there. bmask_t = bmask_b = None if W > 0 and str(getattr(self.cfg, "attn_impl", "sdpa")) == "flex": try: bmask_t = flex_block_mask(P, W, causal=True, sinks=S, device=x.device) if tmask_bidir is not None: bmask_b = flex_block_mask(P, W, causal=False, sinks=S, device=x.device) except Exception: # torch < 2.5, or no compatible backend. Fall back rather than # fail: the sdpa path is numerically identical, only slower. bmask_t = bmask_b = None for blk in self.blocks: if blk.axis == "time": x = blk( x.reshape(B * C, P, -1), bidirectional_rows=bidir_rows, attn_mask=tmask, attn_mask_bidir=tmask_bidir, block_mask=bmask_t, block_mask_bidir=bmask_b, ).view(B, C, P, -1) else: x = ( blk(x.transpose(1, 2).reshape(B * P, C, -1), attn_mask=vmask) .view(B, P, C, -1) .transpose(1, 2) ) x = self.out_mlp(self.norm(x)) out = self.head(x).view(B, C, P, ps, self.cfg.num_quantiles) return out[:, 0] if squeeze_variates else out # ── losses ─────────────────────────────────────────────────────────────────── def pinball_loss(pred_q, target, levels) -> torch.Tensor: """Mean pinball loss. ``pred_q`` ``(..., num_q)``, ``target`` ``(...)``.""" q = torch.tensor(levels, device=pred_q.device, dtype=pred_q.dtype) err = target.unsqueeze(-1) - pred_q return torch.maximum(q * err, (q - 1.0) * err).mean() def pinball_dense( pred_q: torch.Tensor, target: torch.Tensor, levels, *, horizon_mask: torch.Tensor, obs_mask: torch.Tensor | None = None, lam: float = 0.0, ) -> torch.Tensor: """§4 denser supervision: pinball on the CPM-masked region plus ``lam`` × pinball on the observed region. ``pred_q`` is ``(..., num_q)``, ``target`` and both masks broadcast to ``pred_q.shape[:-1]``. Each term is normalised by its own mask weight, so ``lam`` is a clean relative weight rather than something that drifts with how much CPM happened to mask this batch. ``lam = 0`` reduces exactly to masked-region-only training (upstream). The hazard to keep in view: supervising the observed region is next-patch prediction on visible context, which is a SHORTER-horizon task than the one CPM exists to train, and the long-horizon gap versus classical baselines is Toto's stated top open problem. Read this sweep split by term length; the aggregate will hide the trade. """ q = torch.tensor(levels, device=pred_q.device, dtype=pred_q.dtype) err = target.unsqueeze(-1) - pred_q loss = torch.maximum(q * err, (q - 1.0) * err) # (..., num_q) hm = horizon_mask.to(loss.dtype).unsqueeze(-1) h = (loss * hm).sum() / hm.sum().clamp(min=1.0) if lam == 0.0 or obs_mask is None: return h om = obs_mask.to(loss.dtype).unsqueeze(-1) o = (loss * om).sum() / om.sum().clamp(min=1.0) return h + lam * o