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",
]