Download model.py from tensorlink-dev/yumoto-alpha-18m: direct link, hf CLI and curl.
- Browser
- Download file 64.1 kB
-
https://huggingface.co/tensorlink-dev/yumoto-alpha-18m/resolve/main/model.py
- Command line
-
hf download hf://tensorlink-dev/yumoto-alpha-18m/model.py
-
curl -L -o model.py https://huggingface.co/tensorlink-dev/yumoto-alpha-18m/resolve/main/model.py
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 | |
| 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" | |
| 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" | |
| 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 | |