yumoto-alpha-18m / model.py
tensorlink-dev's picture
Upload folder using huggingface_hub
ed7b8dc verified
Raw History Blame Contribute Delete
64.1 kB
"""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:<m0>,<m1>,..." 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