"""Sinkhorn doubly-stochastic constrained linears (MHC) for tensor heads. This module provides a generic, copy-free re-implementation of two bounded-gain operators that the hardened-training stack uses to keep backward gain in a safe band (about 1.6 for the doubly-stochastic mixer, exactly 1.0 for the additive expert residual). They are INSPIRED BY (not copied from) the reference implementations in the inherited training doctrine and the post-hoc MHC training bench, re-expressed as plain ``nn.Module``s with explicit type annotations and no project-specific coupling. Two classes are exported: * :class:`MHCLinear` -- a square ``nn.Linear`` wrapper whose effective weight is ``mix * (ds @ W) + (1 - mix) * W`` where ``ds`` is a Sinkhorn-Knopp doubly- stochastic projection of a learnable ``ds_weight`` parameter. When ``mix`` approaches 1.0 the operator norm is bounded by the doubly-stochastic mixer (stable training at high LR); when ``mix`` approaches 0.0 the layer falls back to the plain ``W`` it wraps. ``mix`` itself is a learnable scalar so the gradient can dial the constraint on or off per head. Non-square linears fall back to a plain ``Linear`` -- the doubly-stochastic bounding only applies to the square case, which is exactly where every tensor head in this module is designed to live. * :class:`MHCExpert` -- a bounded additive residual expert built from two :class:`MHCLinear` projections. ``delta = tanh(MHCLinear(x))`` is bounded to ``[-1, 1]`` and ``alpha = sigmoid(MHCLinear(x))`` is bounded to ``[0, 1]``; the returned residual is ``alpha * delta`` -- bounded in ``[-1, 1]`` by construction regardless of input magnitude, while the Sinkhorn mixers keep the residual transport well conditioned. An :class:`torch.nn.RMSNorm` precedes the projections to keep the input scale well-conditioned. The doubly-stochastic property is produced by :meth:`MHCLinear._sinkhorn`: a ``softplus`` non-negativity projection (randn init can produce negatives, which would break Sinkhorn convergence) followed by ``sinkhorn_iters`` alternating row / column normalizations. After convergence both row sums and column sums are approximately 1 (within the ``1e-8`` clamp floor). """ from __future__ import annotations from typing import Final import torch import torch.nn.functional as F from torch import Tensor, nn MHC_LINEAR_TENSOR_SCHEMA = "nnf.resynthesis.mhc_linear_tensor.v1" #: Default Sinkhorn-Knopp iteration count. Ten alternating row / column #: normalizations bring a uniformly-initialized ``[H,H]`` matrix to doubly- #: stochastic within float32 precision (row / col sums within ``~1e-5`` of 1). #: Higher counts buy more precision at linear cost; lower counts leave a small #: residual imbalance that the learnable ``mix`` parameter can compensate for. DEFAULT_SINKHORN_ITERS: Final[int] = 10 #: Default initial value of the learnable ``mix`` scalar in :class:`MHCLinear`. #: ``0.9`` starts the head strongly constrained (bounded-gain regime); gradient #: descent can pull it toward 0.0 to fall back to the plain wrapped weight or #: toward 1.0 to fully apply the doubly-stochastic mixer. DEFAULT_MIX: Final[float] = 0.9 #: Numerical floor for Sinkhorn normalizations -- keeps divisions finite for a #: degenerate all-zero column without affecting the converged value for any #: non-degenerate init. Matches the reference implementation's clamp. _SINKHORN_EPS: Final[float] = 1e-8 _MHC_WEIGHT_SEED: Final[int] = 0x4D484357 _MHC_SINKHORN_SEED: Final[int] = 0x4D484344 def _deterministic_normal_parameter_t( size: int, *, seed: int, dtype: torch.dtype, device: torch.device | None, ) -> Tensor: """Return a meta-safe MHC seed independent of ambient RNG history.""" value_t = torch.empty(size, size, dtype=dtype, device=device) if value_t.device.type == "meta": return value_t generator = torch.Generator(device=value_t.device) generator.manual_seed(seed + size) return value_t.normal_(mean=0.0, std=0.02, generator=generator) class MHCLinear(nn.Module): """Sinkhorn doubly-stochastic constrained square linear. Wraps a square ``nn.Linear`` (``in_features == out_features``) so that its effective weight is a learnable blend of the plain weight ``W`` and the doubly-stochastic-mixed weight ``ds @ W``: effective_W = mix * (ds @ W) + (1 - mix) * W where ``ds`` is a doubly-stochastic matrix produced by Sinkhorn-Knopp projection of a learnable ``ds_weight`` parameter. Because ``ds`` has bounded operator norm (its rows and columns each sum to 1), the mixed weight has bounded operator norm, which keeps the backward gain of the layer in a safe band and enables stable training at higher learning rates. The ``mix`` scalar is itself learnable (init :data:`DEFAULT_MIX`), so the optimizer can dial the constraint per head: ``mix -> 0`` recovers the plain ``W`` (unconstrained), ``mix -> 1`` fully applies the doubly-stochastic mixer. Non-square linears fall back to a plain ``Linear`` -- the bounding only applies to the square case, and every tensor head designed to use this wrapper is square, so the fallback is a hard constraint rather than a silent skip. Args: size: the square dimension (``in_features == out_features == size``). Must be positive. sinkhorn_iters: number of alternating row / column normalizations in the Sinkhorn-Knopp projection (default :data:`DEFAULT_SINKHORN_ITERS`). mix_init: initial value of the learnable ``mix`` scalar (default :data:`DEFAULT_MIX`). dtype: torch dtype for the parameters. device: torch device for the parameters. Example: >>> import torch >>> from resynthesis.mhc_linear_tensor import MHCLinear >>> head = MHCLinear(size=4) >>> x = torch.randn(8, 4) >>> y = head(x) # bounded-gain forward >>> y.sum().backward() # gradient flows through mix, ds, weight, bias >>> head.mix.item() # learnable scalar, init 0.9 0.9 """ # Class-level annotations make mypy strict happy: nn.Parameter assignments # are otherwise typed as Tensor | nn.Parameter and the attribute access in # forward needs a concrete Tensor type. weight: Tensor bias: Tensor ds_weight: Tensor mix: Tensor def __init__( self, size: int, *, sinkhorn_iters: int = DEFAULT_SINKHORN_ITERS, mix_init: float = DEFAULT_MIX, dtype: torch.dtype = torch.float32, device: torch.device | None = None, ) -> None: super().__init__() if size <= 0: raise ValueError(f"size must be positive, got {size}") if sinkhorn_iters < 1: raise ValueError( f"sinkhorn_iters must be at least 1, got {sinkhorn_iters}" ) self.size = int(size) self._iters = int(sinkhorn_iters) # Plain wrapped linear weight (square). Small randn init keeps the # operator norm modest before the mixer even applies. self.weight = nn.Parameter( _deterministic_normal_parameter_t( self.size, seed=_MHC_WEIGHT_SEED, dtype=dtype, device=device, ) ) self.bias = nn.Parameter(torch.zeros(self.size, dtype=dtype, device=device)) # Learnable doubly-stochastic source. softplus + Sinkhorn below maps # this to a non-negative doubly-stochastic matrix. self.ds_weight = nn.Parameter( _deterministic_normal_parameter_t( self.size, seed=_MHC_SINKHORN_SEED, dtype=dtype, device=device, ) ) # Learnable blend in [0, 1] -- sigmoid keeps it bounded so the head # cannot drift outside the [plain-W, ds-mixed-W] axis. self.mix = nn.Parameter( torch.tensor(float(mix_init), dtype=dtype, device=device) ) # -- Sinkhorn-Knopp doubly-stochastic projection -------------------- def _sinkhorn(self, w: Tensor) -> Tensor: """Project ``w`` to a doubly-stochastic matrix via Sinkhorn-Knopp. Args: w: ``[size, size]`` source matrix (any sign). Returns: ``[size, size]`` non-negative matrix whose row sums and column sums are each approximately 1 (within :data:`_SINKHORN_EPS`). The ``softplus`` first step guarantees non-negativity, which Sinkhorn requires to converge. """ # randn init can produce negatives -> apply softplus before normalizing. # softplus is smooth and strictly positive, which keeps gradients # flowing everywhere (unlike relu, which would zero half the entries). ds = F.softplus(w) for _ in range(self._iters): ds = ds / ds.sum(dim=0, keepdim=True).clamp_min(_SINKHORN_EPS) ds = ds / ds.sum(dim=1, keepdim=True).clamp_min(_SINKHORN_EPS) return ds def doubly_stochastic(self) -> Tensor: """The current doubly-stochastic mixer (for inspection / tests).""" return self._sinkhorn(self.ds_weight) def effective_weight(self) -> Tensor: """The current effective weight ``mix * (ds @ W) + (1 - mix) * W``.""" ds = self.doubly_stochastic() mix = torch.sigmoid(self.mix) return mix * (ds @ self.weight) + (1.0 - mix) * self.weight # -- forward -------------------------------------------------------- def forward(self, x: Tensor) -> Tensor: """Apply the bounded-gain linear: ``effective_weight @ x + bias``. The matmul is factored as ``W`` first then the doubly-stochastic mixer to avoid materializing the full ``[size, size]`` effective weight as a temporary during forward -- the same memory-friendly factoring the reference stack uses. ``x @ W.T`` is the plain linear, then the mixer is applied to the result. """ # mix in [0, 1] via sigmoid so the scalar stays bounded. mix = torch.sigmoid(self.mix) ds = self.doubly_stochastic() base = F.linear(x, self.weight) # ``F.linear(base, ds) == base @ ds.T``. Since # ``base == x @ W.T``, this is exactly # ``x @ W.T @ ds.T == x @ (ds @ W).T`` and therefore matches # ``effective_weight()``. Passing ``ds.T`` here would instead apply # ``W.T @ ds`` and silently train a different operator. mixed = mix * F.linear(base, ds) + (1.0 - mix) * base return mixed + self.bias class MHCExpert(nn.Module): """Bounded-residual additive expert built from two :class:`MHCLinear` heads. Produces a bounded residual ``alpha * delta`` where ``delta = tanh(...)`` is in ``[-1, 1]`` and ``alpha = sigmoid(...)`` is in ``[0, 1]`` -- so the returned residual is in ``[-1, 1]`` by construction regardless of input magnitude. Both projections are :class:`MHCLinear` (Sinkhorn-bounded), so the backward gain of the expert is bounded by the doubly-stochastic mixers and the expert trains stably at any learning rate. The signal fed to both projections is the RMS-normalized input (a single ``hidden`` tensor). Because the two projections differ only in their activation (tanh for the delta head, sigmoid for the alpha head), they share their input but learn independent bounded-gain weights. Args: size: the square dimension of both MHC heads (the input feature size). Must be positive. sinkhorn_iters: forwarded to both :class:`MHCLinear` heads. mix_init: forwarded to both :class:`MHCLinear` heads. dtype: torch dtype for the parameters. device: torch device for the parameters. Example: >>> import torch >>> from resynthesis.mhc_linear_tensor import MHCExpert >>> expert = MHCExpert(size=4) >>> hidden = torch.randn(2, 3, 4) # [batch, seq, hidden] >>> residual = expert(hidden) # bounded in [-1, 1] >>> residual.shape torch.Size([2, 3, 4]) >>> residual.abs().max().item() <= 1.0 True """ # Class-level annotations for mypy strict. delta_head: MHCLinear alpha_head: MHCLinear def __init__( self, size: int, *, sinkhorn_iters: int = DEFAULT_SINKHORN_ITERS, mix_init: float = DEFAULT_MIX, dtype: torch.dtype = torch.float32, device: torch.device | None = None, ) -> None: super().__init__() if size <= 0: raise ValueError(f"size must be positive, got {size}") self.size = int(size) # RMSNorm precedes the projections to keep the input scale # well-conditioned (so tanh does not saturate and sigmoid stays in its # linear region). PyTorch's nn.RMSNorm is the canonical impl; we # forward dtype/device so the affine weight matches the heads (avoids # an internal upcast warning on the layer_norm kernel). self.norm: nn.RMSNorm = nn.RMSNorm( self.size, dtype=dtype, device=device ) self.delta_head = MHCLinear( self.size, sinkhorn_iters=sinkhorn_iters, mix_init=mix_init, dtype=dtype, device=device, ) # Alpha head projects to a scalar per token -- implemented as a square # ``size`` head whose output we reduce to the last dim. We keep the # head square (size -> size) and take a learned-linear reduction down # to 1 inside forward, so the doubly-stochastic bounding applies # uniformly. This matches the reference expert's "alpha = sigmoid of # a bounded projection" contract. self.alpha_head = MHCLinear( self.size, sinkhorn_iters=sinkhorn_iters, mix_init=mix_init, dtype=dtype, device=device, ) def forward(self, hidden: Tensor) -> Tensor: """Return the bounded residual ``alpha * delta`` (same shape as input). ``hidden`` may be any shape ending in ``size`` (``[size]``, ``[B, size]``, ``[B, S, size]``, ...). The returned tensor has the same shape and is element-wise bounded in ``[-1, 1]``. """ normed = self.norm(hidden) delta = torch.tanh(self.delta_head(normed)) # bounded [-1, 1] alpha_raw = self.alpha_head(normed) # Reduce the alpha projection to a per-token scalar in [0, 1] by # averaging the sigmoided entries along the feature axis. This keeps # the alpha head square (so the doubly-stochastic bounding applies) # while producing a single gating scalar per token as the reference # expert does. alpha = torch.sigmoid(alpha_raw.mean(dim=-1, keepdim=True)) return alpha * delta __all__ = [ "DEFAULT_MIX", "DEFAULT_SINKHORN_ITERS", "MHC_LINEAR_TENSOR_SCHEMA", "MHCExpert", "MHCLinear", ]