File size: 15,158 Bytes
919fd68 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 | """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",
]
|