"""Cross-loop Transition Adapter (CTA) from LoopQ Equation (7).""" from __future__ import annotations import math from collections.abc import Mapping from typing import Any import torch from torch import nn DEFAULT_CTA_RANK = 8 DEFAULT_RMSNORM_EPSILON = 1e-6 OURO_TRANSITION_COUNT = 3 HUGINN_TRANSITION_COUNT = 31 class CrossLoopTransitionAdapter(nn.Module): """Apply LoopQ CTA only at a true cross-loop transition. ``U`` and ``V`` are shared across all transitions. The affine vectors ``a_t``, ``b_t`` and low-rank gate ``eta_t`` are transition-dependent. Identity initialization is exact because ``a=1``, ``b=0`` and ``eta=0``. """ FORMAT_VERSION = 1 def __init__( self, hidden_size: int, transition_count: int, *, rank: int = DEFAULT_CTA_RANK, rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, dtype: torch.dtype = torch.float32, ) -> None: super().__init__() if hidden_size <= 0: raise ValueError("hidden_size must be positive") if transition_count <= 0: raise ValueError("transition_count must be positive") if rank <= 0 or rank > hidden_size: raise ValueError("rank must be positive and no larger than hidden_size") if not math.isfinite(rmsnorm_epsilon) or rmsnorm_epsilon <= 0: raise ValueError("rmsnorm_epsilon must be finite and positive") self.enabled = True self.hidden_size = int(hidden_size) self.transition_count = int(transition_count) self.rank = int(rank) self.rmsnorm_epsilon = float(rmsnorm_epsilon) # Deterministic full-rank column initialization. eta=0 keeps the CTA # exactly identity while leaving a nonzero gradient path to its gates. basis = torch.eye(hidden_size, dtype=dtype)[:, :rank] self.U = nn.Parameter(basis.clone()) self.V = nn.Parameter(basis.clone()) self.a = nn.Parameter(torch.ones(transition_count, hidden_size, dtype=dtype)) self.b = nn.Parameter(torch.zeros(transition_count, hidden_size, dtype=dtype)) self.eta = nn.Parameter(torch.zeros(transition_count, rank, dtype=dtype)) @classmethod def for_ouro( cls, hidden_size: int, *, rank: int = DEFAULT_CTA_RANK, rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, ) -> "CrossLoopTransitionAdapter": return cls( hidden_size, OURO_TRANSITION_COUNT, rank=rank, rmsnorm_epsilon=rmsnorm_epsilon, ) @classmethod def for_huginn( cls, hidden_size: int, *, rank: int = DEFAULT_CTA_RANK, rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, ) -> "CrossLoopTransitionAdapter": return cls( hidden_size, HUGINN_TRANSITION_COUNT, rank=rank, rmsnorm_epsilon=rmsnorm_epsilon, ) def _rmsnorm(self, hidden_state: torch.Tensor) -> torch.Tensor: work = hidden_state.to(torch.float32) normalized = work * torch.rsqrt(work.square().mean(dim=-1, keepdim=True) + self.rmsnorm_epsilon) return normalized.to(hidden_state.dtype) def forward(self, hidden_state: torch.Tensor, transition_index: int) -> torch.Tensor: if hidden_state.ndim == 0 or hidden_state.shape[-1] != self.hidden_size: raise ValueError( f"hidden_state last dimension must be {self.hidden_size}, " f"got {tuple(hidden_state.shape)}" ) if not 0 <= transition_index < self.transition_count: raise IndexError( f"transition_index must be in [0, {self.transition_count}), " f"got {transition_index}" ) if not self.enabled: return hidden_state normalized = self._rmsnorm(hidden_state) a = self.a[transition_index].to(hidden_state) b = self.b[transition_index].to(hidden_state) eta = self.eta[transition_index].to(hidden_state) U = self.U.to(hidden_state) V = self.V.to(hidden_state) affine = (a - 1) * normalized + b low_rank = ((normalized @ V) * eta) @ U.T return hidden_state + affine + low_rank def export_state(self) -> dict[str, Any]: return { "format": "loopq_cta", "format_version": self.FORMAT_VERSION, "hidden_size": self.hidden_size, "transition_count": self.transition_count, "rank": self.rank, "enabled": self.enabled, "rmsnorm_epsilon": self.rmsnorm_epsilon, "state_dict": {name: value.detach().cpu() for name, value in self.state_dict().items()}, } @classmethod def from_export_state(cls, state: Mapping[str, Any]) -> "CrossLoopTransitionAdapter": required = { "format", "format_version", "hidden_size", "transition_count", "rank", "rmsnorm_epsilon", "state_dict", } missing = required.difference(state) if missing: raise ValueError(f"CTA export is missing fields: {sorted(missing)}") if state["format"] != "loopq_cta" or state["format_version"] != cls.FORMAT_VERSION: raise ValueError("unsupported CTA export format or version") adapter = cls( int(state["hidden_size"]), int(state["transition_count"]), rank=int(state["rank"]), rmsnorm_epsilon=float(state["rmsnorm_epsilon"]), dtype=state["state_dict"]["U"].dtype, ) adapter.load_state_dict(state["state_dict"]) adapter.enabled = bool(state.get("enabled", True)) if not adapter.enabled: adapter.requires_grad_(False) return adapter