| """State-space models, 1-D and N-D under one name. |
| |
| ``td.S4(d_model, n_layers)`` with no lattice is a sequence model; give it a |
| lattice and it is S4ND. Same for ``td.S4D`` and ``td.Mamba``. The explicit N-D |
| names — ``td.S4ND``, ``td.S4DND``, ``td.MambaND`` — are the same classes with |
| ``dim`` and the lattice made mandatory: taking the N-D name means declaring |
| what N is, the declaration is checked, and ``dim=1`` is refused outright — |
| one spatial axis is the 1-D model, and code reading "S4ND" must not be |
| running S4. Pass ``lattice=...`` or just ``shape=(32, 32)`` and the lattice |
| is constructed for you. |
| |
| **The default mixers are the original authors' code**, shipped verbatim in |
| ``torch_dimensions._vendor`` and byte-verified against their repositories: |
| S4/S4D construct upstream's real ``S4Block`` through their own hydra |
| registry, Mamba runs the reference block with the authors' own selective |
| scan. Their dependencies (einops, numpy, scipy, hydra-core, omegaconf) are |
| installed on first use — never at import, never for ``portable=True``. |
| |
| ``portable=True`` selects our pure-torch implementations in |
| :mod:`torch_dimensions.mixers.ssm` instead: no dependencies beyond torch, |
| verified to agree with the originals (the S4D kernel bitwise). The flag is |
| recorded in the model's config, so checkpoints rebuild what was actually |
| trained; checkpoints written before this flag existed rebuild portable, which |
| is what they were. |
| |
| How the axes are composed stays ``nd_method``'s business (default |
| :func:`~torch_dimensions.axial_scan`), exactly as for the RNN family. |
| """ |
|
|
| from __future__ import annotations |
|
|
| from collections.abc import Sequence |
|
|
| import torch |
| import torch.nn as nn |
|
|
| from torch_dimensions.lattice import Lattice |
| from torch_dimensions.mixers.ssm import MambaMixer, S4DMixer, S4Mixer |
| from torch_dimensions.mixers.upstream import ( |
| Mamba3Mixer, |
| UpstreamMamba2Mixer, |
| UpstreamMambaMixer, |
| UpstreamS4DMixer, |
| UpstreamS4Mixer, |
| ) |
| from torch_dimensions.models.base import LatticeModel |
|
|
| __all__ = [ |
| "S4", |
| "S4D", |
| "S4DND", |
| "S4ND", |
| "Mamba", |
| "Mamba2", |
| "Mamba2ND", |
| "Mamba3", |
| "Mamba3ND", |
| "MambaND", |
| ] |
|
|
|
|
| def _pick( |
| portable: bool, portable_cls: type[nn.Module], upstream_cls: type[nn.Module], kw: dict |
| ) -> type[nn.Module]: |
| """The flag chooses the implementation; an explicit mixer= makes it moot, |
| and combining the two is a contradiction refused rather than resolved.""" |
| if portable and kw.get("mixer") is not None: |
| raise ValueError("pass either portable=True or mixer=..., not both") |
| return portable_cls if portable else upstream_cls |
|
|
|
|
| class S4(LatticeModel): |
| """The full S4 (DPLR: diagonal plus low-rank) over a sequence or lattice. |
| |
| Args: |
| d_model: feature width. |
| n_layers: sweeps; with a lattice, layers cycle through its axes. |
| lattice: omit for an ordinary 1-D sequence model. |
| d_state: full SSM state size (even; stored as conjugate pairs). |
| |
| The kernel carries the rank-1 HiPPO-LegS correction that distinguishes S4 |
| from its diagonal approximation :class:`S4D`. By default this is |
| upstream's real ``S4Block(mode='dplr')``; ``portable=True`` selects our |
| pure-torch kernel instead. |
| """ |
|
|
| _mixer: type[nn.Module] = S4Mixer |
|
|
| def __init__( |
| self, |
| d_model: int, |
| n_layers: int = 1, |
| lattice=None, |
| *, |
| d_state: int = 64, |
| portable: bool = False, |
| **kw, |
| ): |
| self._mixer = _pick(portable, S4Mixer, UpstreamS4Mixer, kw) |
| mixer_kwargs = {"d_state": d_state, **kw.pop("mixer_kwargs", {})} |
| super().__init__(d_model, n_layers, lattice, mixer_kwargs=mixer_kwargs, **kw) |
| self.config["portable"] = portable |
|
|
|
|
| class S4D(LatticeModel): |
| """Diagonal state-space model (S4D) over a sequence or an N-D lattice. |
| |
| Args: |
| d_model: feature width. |
| n_layers: sweeps; with a lattice, layers cycle through its axes. |
| lattice: omit for an ordinary 1-D sequence model. |
| d_state: state dimension of the diagonal SSM (even; conjugate pairs). |
| |
| Extra mixer options (``dt_min``, ``dt_max``) go in ``mixer_kwargs``. By |
| default this is upstream's real ``S4Block(mode='diag')``; |
| ``portable=True`` selects our pure-torch kernel instead. |
| """ |
|
|
| _mixer: type[nn.Module] = S4DMixer |
|
|
| def __init__( |
| self, |
| d_model: int, |
| n_layers: int = 1, |
| lattice=None, |
| *, |
| d_state: int = 64, |
| portable: bool = False, |
| **kw, |
| ): |
| self._mixer = _pick(portable, S4DMixer, UpstreamS4DMixer, kw) |
| mixer_kwargs = {"d_state": d_state, **kw.pop("mixer_kwargs", {})} |
| super().__init__(d_model, n_layers, lattice, mixer_kwargs=mixer_kwargs, **kw) |
| self.config["portable"] = portable |
|
|
|
|
| class Mamba(LatticeModel): |
| """Mamba (selective SSM) over a sequence or an N-D lattice. |
| |
| With a lattice this is the Mamba-ND construction: each layer runs the |
| selective scan along one axis, and the :class:`~torch_dimensions.ScanPlan` |
| decides which axis and direction — including the paired schedule of the |
| official Mamba-ND implementation via ``ScanPlan.paired``. |
| |
| Args: |
| d_state: SSM state size per channel. |
| d_conv: width of the causal depthwise convolution. |
| expand: inner width multiplier. |
| |
| By default the layer is the authors' reference ``Mamba`` block (their |
| selective scan, end to end); ``portable=True`` selects our pure-torch |
| mixer instead. |
| |
| ``version=2`` runs the authors' **Mamba-2** block (the SSD formulation: |
| multi-head, gated RMSNorm) and ``version=3`` their **Mamba-3** block |
| (rotary state, trapezoidal discretization) — the same objects as |
| :class:`Mamba2` and :class:`Mamba3`, which are simply the spellings that |
| put the version in the name. |
| |
| Neither has a ``portable`` build, for different reasons. Mamba-2 already |
| runs the authors' own reference SSD off GPU, so a reimplementation would |
| add nothing but a second thing to be wrong. Mamba-3 has no upstream |
| reference at all: its scan is Triton-only, so off GPU the recurrence is |
| computed by our transcription of it — see |
| :class:`~torch_dimensions.mixers.Mamba3Mixer`, which is named without |
| ``Upstream`` precisely because those numbers are ours. |
| """ |
|
|
| _mixer: type[nn.Module] = MambaMixer |
|
|
| def __init__( |
| self, |
| d_model: int, |
| n_layers: int = 1, |
| lattice=None, |
| *, |
| d_state: int | None = None, |
| d_conv: int = 4, |
| expand: int = 2, |
| portable: bool = False, |
| version: int = 1, |
| **kw, |
| ): |
| if version not in (1, 2, 3): |
| raise ValueError(f"Mamba version must be 1, 2 or 3; got {version}") |
| if version in (2, 3) and portable: |
| raise ValueError( |
| f"there is no portable build of Mamba-{version}: off GPU it already runs " |
| + ( |
| "the authors' own reference SSD implementation" |
| if version == 2 |
| else "their block around a transcribed scan (mixers/mamba3_compat.py)" |
| ) |
| + ". Use version=1 for the portable selective scan." |
| ) |
| if version == 3: |
| self._mixer = _pick(False, MambaMixer, Mamba3Mixer, kw) |
| |
| defaults: dict = {"d_state": 128 if d_state is None else d_state} |
| elif version == 2: |
| self._mixer = _pick(False, MambaMixer, UpstreamMamba2Mixer, kw) |
| defaults = {"d_state": 128 if d_state is None else d_state, "d_conv": d_conv} |
| else: |
| self._mixer = _pick(portable, MambaMixer, UpstreamMambaMixer, kw) |
| defaults = {"d_state": 16 if d_state is None else d_state, "d_conv": d_conv} |
| mixer_kwargs = {**defaults, "expand": expand, **kw.pop("mixer_kwargs", {})} |
| super().__init__(d_model, n_layers, lattice, mixer_kwargs=mixer_kwargs, **kw) |
| self.config["portable"] = portable |
| self.config["version"] = version |
|
|
|
|
| def _nd_lattice( |
| cls_name: str, |
| lattice: Lattice | None, |
| shape: Sequence[int] | None, |
| names: Sequence[str] | None, |
| valid: torch.Tensor | None, |
| time: bool, |
| dim: int | None, |
| ) -> Lattice: |
| """Resolve the N-D classes' lattice sugar, refusing the ambiguous cases.""" |
| |
| |
| |
| |
| if dim is None: |
| raise ValueError( |
| f"{cls_name} requires `dim` — declare the number of spatial axes, " |
| f"e.g. td.{cls_name}(64, 8, dim=2, shape=(32, 32))" |
| ) |
| base = cls_name.removesuffix("ND") |
| if dim == 1: |
| |
| |
| |
| raise ValueError( |
| f"dim=1 is not {cls_name} — one spatial axis is just {base}. " |
| f"Use td.{base}(..., lattice=...) so the code says what is actually running." |
| ) |
| if dim < 1: |
| raise ValueError(f"the N-D names need dim >= 2; got {dim}") |
| if lattice is None: |
| if shape is None: |
| raise ValueError( |
| f"{cls_name} needs a lattice — pass `lattice=...` or let it build one: " |
| f"td.{cls_name}(64, 8, shape=(32, 32))" |
| ) |
| lattice = Lattice( |
| shape=tuple(shape), names=tuple(names) if names else None, valid=valid, time=time |
| ) |
| elif shape is not None or names is not None or valid is not None: |
| raise ValueError("pass either `lattice` or `shape`/`names`/`valid`, not both") |
| if lattice.rank < 1: |
| raise ValueError( |
| f"{cls_name} is the N-D name and this lattice has no spatial axes; " |
| f"for a plain sequence use td.{cls_name.removesuffix('ND')}" |
| ) |
| if lattice.rank != dim: |
| raise ValueError(f"dim={dim}, but the lattice has {lattice.rank} spatial axes") |
| return lattice |
|
|
|
|
| def _nd_variant(base: type[LatticeModel], cls_name: str) -> type[LatticeModel]: |
| class ND(base): |
| def __init__( |
| self, |
| d_model: int, |
| n_layers: int = 1, |
| lattice: Lattice | None = None, |
| *, |
| shape: Sequence[int] | None = None, |
| names: Sequence[str] | None = None, |
| valid: torch.Tensor | None = None, |
| time: bool = True, |
| dim: int | None = None, |
| **kw, |
| ): |
| lat = _nd_lattice(cls_name, lattice, shape, names, valid, time, dim) |
| super().__init__(d_model, n_layers, lat, **kw) |
| |
| |
| self.config["dim"] = lat.rank |
|
|
| ND.__name__ = ND.__qualname__ = cls_name |
| ND.__doc__ = ( |
| f"{base.__name__} with `dim` and a lattice mandatory — the explicit N-D name.\n\n" |
| f" td.{cls_name}(64, 8, dim=2, shape=(32, 32)) # builds the lattice\n" |
| f" td.{cls_name}(64, 8, dim=2, lattice=my_lattice) # or bring your own\n\n" |
| "Taking the N-D name means declaring what N is; ``dim`` is checked\n" |
| "against the lattice's spatial rank. ``time=True`` by default.\n" |
| f"Identical to ``td.{base.__name__}`` with a lattice in every other way." |
| ) |
| return ND |
|
|
|
|
| class Mamba2(Mamba): |
| """Mamba-2 (the SSD formulation) — ``td.Mamba(..., version=2)`` by name. |
| |
| Identical in every way to passing ``version=2``; both spellings build the |
| same model and record the same config, so a checkpoint written by one |
| rebuilds under the other. Which reads better is the caller's choice. |
| |
| Off GPU the chunked scan is computed by the authors' own reference |
| implementation (``ssd_minimal.py``, "the same as Listing 1 from the |
| paper"); with Triton and CUDA present the fused kernels are used exactly |
| as upstream intends. |
| """ |
|
|
| def __init__(self, d_model: int, n_layers: int = 1, lattice=None, **kw): |
| if kw.pop("version", 2) != 2: |
| raise ValueError("td.Mamba2 is version 2; use td.Mamba(version=...) to choose") |
| super().__init__(d_model, n_layers, lattice, version=2, **kw) |
|
|
|
|
| class Mamba3(Mamba): |
| """Mamba-3 — ``td.Mamba(..., version=3)`` by name. |
| |
| Identical in every way to passing ``version=3``; both spellings build the |
| same model and record the same config, so a checkpoint written by one |
| rebuilds under the other. |
| |
| The block is the authors', verbatim — rotary state, trapezoidal |
| discretization, heavy-tail ``A``. The *scan* is theirs only on CUDA: it |
| ships upstream as Triton alone, so elsewhere it is computed by our |
| transcription of the same recurrence. See |
| :class:`~torch_dimensions.mixers.Mamba3Mixer` for exactly what that does |
| and does not establish. |
| """ |
|
|
| def __init__(self, d_model: int, n_layers: int = 1, lattice=None, **kw): |
| if kw.pop("version", 3) != 3: |
| raise ValueError("td.Mamba3 is version 3; use td.Mamba(version=...) to choose") |
| super().__init__(d_model, n_layers, lattice, version=3, **kw) |
|
|
|
|
| S4ND = _nd_variant(S4, "S4ND") |
| S4DND = _nd_variant(S4D, "S4DND") |
| MambaND = _nd_variant(Mamba, "MambaND") |
| Mamba2ND = _nd_variant(Mamba2, "Mamba2ND") |
| Mamba3ND = _nd_variant(Mamba3, "Mamba3ND") |
|
|