amoe-lora / src /amoe /core /adapter.py
AbstractPhil's picture
amoe-lora 0.1.0: working-state framework — attach/toggle/detach invariant-tested (toggle law + bit-exact detach), checkpoint v1 + legacy import verified against shipped campaign artifacts, reference-grade train/align with guards, DDP-aware, honesty diagnostics first-class
906715b verified
Raw
History Blame Contribute Delete
2.4 kB
"""RelayPatchwork adapter + the hybrid-safe block wrapper.
The certified geometry (16 slots x D=4 through a 64-atom aleph address,
squared-ReLU patch head, sigmoid gate initialized at -3): ~261k params
at d=1024. Output head zero-initialized so a fresh anchor is inert
(modulo the documented LayerNorm-bias leak).
BlockWithAdapter carries an `enabled` switch: when False, the forward
returns the block output untouched — the single-anchor half of the
toggle law.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from .address import AlephAddress
@dataclass
class AdapterSpec:
n_slots: int = 16
K: int = 64
D: int = 4
tau: float = 0.1
hidden: int = 178
gate_init: float = -3.0
zero_init_head: bool = True
class SquaredReLU(nn.Module):
def forward(self, x):
return F.relu(x) ** 2
class RelayPatchwork(nn.Module):
def __init__(self, d: int, spec: AdapterSpec | None = None):
super().__init__()
s = spec or AdapterSpec()
self.spec = s
self.n_slots = s.n_slots
self.proj = nn.Linear(d, s.n_slots * s.D, bias=False)
nn.init.orthogonal_(self.proj.weight)
self.addr = AlephAddress(s.K, s.D, s.tau)
self.consume = nn.Sequential(
nn.Linear(s.n_slots * s.D, s.hidden), SquaredReLU(),
nn.LayerNorm(s.hidden), nn.Linear(s.hidden, d))
if s.zero_init_head:
nn.init.zeros_(self.consume[-1].weight)
self.gate = nn.Parameter(torch.tensor(float(s.gate_init)))
def forward(self, x):
B, n, _ = x.shape
slots = self.proj(x).view(B, n, self.n_slots, self.spec.D)
feats = self.addr.m_hat(slots).reshape(B, n, -1)
return x + torch.sigmoid(self.gate) * self.consume(feats)
class BlockWithAdapter(nn.Module):
"""Wraps one decoder block; hybrid-safe (tuple or tensor output)."""
def __init__(self, block: nn.Module, adapter: RelayPatchwork):
super().__init__()
self.block = block
self.adapter = adapter
self.enabled = True
def forward(self, *args, **kwargs):
out = self.block(*args, **kwargs)
if not self.enabled:
return out
if isinstance(out, tuple):
return (self.adapter(out[0]),) + out[1:]
return self.adapter(out)