File size: 5,891 Bytes
9118991 | 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 | """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
|