"""Fixed-budget lexical residual memory for graph-quotient language models.""" from __future__ import annotations from dataclasses import dataclass from enum import IntEnum import math import torch from torch.nn import functional as F class ResidualNodeType(IntEnum): """Retrieval/eviction types; these are never semantic-role addresses.""" LEXICAL_FORM = 0 RARE_ENTITY_SURFACE = 1 NUMBER_OR_DATE = 2 QUOTED_SPAN = 3 UNRESOLVED_SPAN = 4 DISCOURSE_STYLE = 5 LOCAL_SUMMARY_RESIDUAL = 6 @dataclass(frozen=True, slots=True) class ResidualMemory: """One immutable sparse memory shared by all queries in a sequence.""" keys: torch.Tensor values: torch.Tensor token_positions: torch.Tensor segment_ids: torch.Tensor node_types: torch.Tensor utility: torch.Tensor budget: int def __post_init__(self) -> None: if self.keys.ndim != 2 or self.values.ndim != 2: raise ValueError("residual keys and values must be [slots, dimension]") if self.keys.shape != self.values.shape: raise ValueError("residual keys and values must have identical shapes") slots = self.keys.shape[0] for name in ("token_positions", "segment_ids", "node_types", "utility"): value = getattr(self, name) if value.ndim != 1 or value.shape[0] != slots: raise ValueError(f"{name} must contain one value per residual slot") if self.budget <= 0 or slots > self.budget: raise ValueError("residual memory exceeds its positive fixed budget") if slots and not bool((self.token_positions >= 0).all()): raise ValueError("residual token positions must be non-negative") if slots and len(set(int(value) for value in self.token_positions.tolist())) != slots: raise ValueError("residual token positions must be unique") @property def slots(self) -> int: return self.keys.shape[0] @dataclass(frozen=True, slots=True) class SparseResidualRead: values: torch.Tensor selected_slots: torch.Tensor selected_count: torch.Tensor @dataclass(frozen=True, slots=True) class BoundedFusionOutput: hidden: torch.Tensor graph_update: torch.Tensor residual_update: torch.Tensor combined_update: torch.Tensor relative_update_rms: torch.Tensor def reconstruction_residual( reference_hidden: torch.Tensor, graph_quotient_hidden: torch.Tensor, ) -> torch.Tensor: """Return a stopped-gradient reference-minus-quotient residual.""" if reference_hidden.shape != graph_quotient_hidden.shape: raise ValueError("reference and graph-quotient hidden states must match") return (reference_hidden.detach() - graph_quotient_hidden.detach()).detach() def residual_utility( residual: torch.Tensor, future_relevance: torch.Tensor, *, epsilon: float = 1e-8, ) -> torch.Tensor: """Score items by reconstruction failure times non-negative relevance.""" if residual.ndim != 2 or future_relevance.shape != residual.shape[:1]: raise ValueError("residual must be [tokens, dim] with one relevance per token") failure = residual.float().pow(2).mean(dim=-1).add(epsilon).sqrt() return failure * future_relevance.float().clamp_min(0) def select_residual_memory( residual: torch.Tensor, future_relevance: torch.Tensor, *, budget: int, segment_size: int, retrieval_keys: torch.Tensor | None = None, token_positions: torch.Tensor | None = None, node_types: torch.Tensor | None = None, maximum_segment_fraction: float = 0.25, maximum_type_fraction: float = 0.75, ) -> ResidualMemory: """Select deterministic high-utility slots under segment/type diversity caps.""" if residual.ndim != 2: raise ValueError("residual must be [tokens, dimension]") if budget <= 0 or segment_size <= 0: raise ValueError("budget and segment_size must be positive") tokens = residual.shape[0] device = residual.device if token_positions is None: token_positions = torch.arange(tokens, device=device, dtype=torch.long) if node_types is None: node_types = torch.full( (tokens,), int(ResidualNodeType.LEXICAL_FORM), device=device, dtype=torch.long, ) if token_positions.shape != (tokens,) or node_types.shape != (tokens,): raise ValueError("token positions and node types must align with residual tokens") if retrieval_keys is not None and retrieval_keys.shape != residual.shape: raise ValueError("retrieval keys must match residual shape") if not 0 < maximum_segment_fraction <= 1 or not 0 < maximum_type_fraction <= 1: raise ValueError("diversity fractions must lie in (0, 1]") utility = residual_utility(residual, future_relevance) segment_ids = torch.div(token_positions, segment_size, rounding_mode="floor") utility_values = utility.detach().cpu().tolist() position_values = token_positions.detach().cpu().tolist() segment_values = segment_ids.detach().cpu().tolist() type_values = node_types.detach().cpu().tolist() ranked = sorted( range(tokens), key=lambda index: (-utility_values[index], position_values[index]), ) segment_cap = max(1, math.ceil(budget * maximum_segment_fraction)) type_cap = max(1, math.ceil(budget * maximum_type_fraction)) segment_counts: dict[int, int] = {} type_counts: dict[int, int] = {} selected: list[int] = [] for index in ranked: if len(selected) == budget: break segment = int(segment_values[index]) node_type = int(type_values[index]) if segment_counts.get(segment, 0) >= segment_cap: continue if type_counts.get(node_type, 0) >= type_cap: continue selected.append(index) segment_counts[segment] = segment_counts.get(segment, 0) + 1 type_counts[node_type] = type_counts.get(node_type, 0) + 1 # Diversity caps can leave capacity unused in degenerate, single-type inputs. if len(selected) < min(tokens, budget): chosen = set(selected) selected.extend(index for index in ranked if index not in chosen) selected = selected[: min(tokens, budget)] indices = torch.tensor(selected, device=device, dtype=torch.long) values = residual.index_select(0, indices) key_source = values if retrieval_keys is None else retrieval_keys.index_select(0, indices) keys = F.normalize(key_source.float(), dim=-1).to(values.dtype) return ResidualMemory( keys=keys, values=values, token_positions=token_positions.index_select(0, indices), segment_ids=segment_ids.index_select(0, indices), node_types=node_types.index_select(0, indices), utility=utility.index_select(0, indices), budget=budget, ) def sparse_residual_read( queries: torch.Tensor, memory: ResidualMemory, *, query_positions: torch.Tensor | None = None, maximum_segments: int = 4, maximum_reads: int = 8, minimum_distance: int = 1, temperature: float = 1.0, ) -> SparseResidualRead: """Perform exact hard segment selection followed by exact hard slot selection.""" if queries.ndim != 2 or queries.shape[1] != memory.values.shape[1]: raise ValueError("queries must be [queries, memory dimension]") if maximum_segments <= 0 or maximum_reads <= 0 or minimum_distance < 1: raise ValueError("selection bounds must be positive") if temperature <= 0: raise ValueError("temperature must be positive") count = queries.shape[0] device = queries.device if query_positions is None: query_positions = torch.arange(count, device=device, dtype=torch.long) if query_positions.shape != (count,): raise ValueError("query positions must align with queries") if memory.slots == 0: empty = torch.empty((count, 0), dtype=torch.long, device=device) return SparseResidualRead(torch.zeros_like(queries), empty, query_positions.new_zeros(count)) normalized_queries = F.normalize(queries.float(), dim=-1) unique_segments = torch.unique(memory.segment_ids, sorted=True) segment_keys = torch.stack([ F.normalize(memory.keys[memory.segment_ids == segment].float().mean(dim=0), dim=0) for segment in unique_segments ]) segment_scores = normalized_queries @ segment_keys.T top_segment_count = min(maximum_segments, unique_segments.numel()) top_segments = unique_segments[segment_scores.topk(top_segment_count, dim=-1).indices] slot_scores = normalized_queries @ memory.keys.float().T slot_segment_allowed = ( memory.segment_ids.view(1, 1, -1) == top_segments.unsqueeze(-1) ).any(dim=1) causal = memory.token_positions.unsqueeze(0) <= ( query_positions.unsqueeze(1) - minimum_distance ) allowed = slot_segment_allowed & causal slot_scores = slot_scores.masked_fill(~allowed, float("-inf")) read_count = min(maximum_reads, memory.slots) top = slot_scores.topk(read_count, dim=-1) valid = torch.isfinite(top.values) safe_scores = top.values.masked_fill(~valid, -1e9) weights = torch.softmax(safe_scores / temperature, dim=-1) * valid weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8) selected_values = memory.values[top.indices] values = torch.einsum("qr,qrd->qd", weights.to(selected_values.dtype), selected_values) values = torch.where(valid.any(dim=-1, keepdim=True), values, torch.zeros_like(values)) indices = top.indices.masked_fill(~valid, -1) return SparseResidualRead(values, indices, valid.sum(dim=-1)) def _rms(value: torch.Tensor, epsilon: float = 1e-8) -> torch.Tensor: return value.float().pow(2).mean(dim=-1, keepdim=True).add(epsilon).sqrt() def _bounded_update( hidden: torch.Tensor, value: torch.Tensor | None, gate: float, cap: float, ) -> torch.Tensor: if value is None or gate == 0 or not bool(value.detach().ne(0).any()): return torch.zeros_like(hidden) if value.shape != hidden.shape: raise ValueError("fusion values must match hidden states") scale = _rms(hidden) / _rms(value) relative_gate = min(abs(float(gate)), cap) return value * scale.to(value.dtype) * math.copysign(relative_gate, float(gate)) def bounded_graph_residual_fusion( hidden: torch.Tensor, *, graph_value: torch.Tensor | None = None, residual_value: torch.Tensor | None = None, graph_gate: float = 0.0, residual_gate: float = 0.0, graph_cap: float = 0.05, residual_cap: float = 0.05, combined_cap: float = 0.08, ) -> BoundedFusionOutput: """Fuse independent channels while enforcing per-channel and total RMS caps.""" if min(graph_cap, residual_cap, combined_cap) < 0: raise ValueError("fusion caps must be non-negative") graph_update = _bounded_update(hidden, graph_value, graph_gate, graph_cap) residual_update = _bounded_update(hidden, residual_value, residual_gate, residual_cap) combined = graph_update + residual_update relative = _rms(combined) / _rms(hidden) scale = (combined_cap / relative).clamp(max=1.0) combined = combined * scale.to(combined.dtype) relative = (_rms(combined) / _rms(hidden)).squeeze(-1) if not bool(combined.detach().ne(0).any()): fused = hidden else: fused = hidden + combined return BoundedFusionOutput(fused, graph_update, residual_update, combined, relative) def graph_residual_orthogonality_loss( graph_value: torch.Tensor, residual_value: torch.Tensor, ) -> torch.Tensor: if graph_value.shape != residual_value.shape: raise ValueError("graph and residual values must match") graph = F.normalize(graph_value.float(), dim=-1) residual = F.normalize(residual_value.float(), dim=-1) return (graph * residual).sum(dim=-1).abs().mean() __all__ = [ "BoundedFusionOutput", "ResidualMemory", "ResidualNodeType", "SparseResidualRead", "bounded_graph_residual_fusion", "graph_residual_orthogonality_loss", "reconstruction_residual", "residual_utility", "select_residual_memory", "sparse_residual_read", ]