""" Memory Sparse Attention (MSA) for PC-SHO-DLM Implements the MSA framework from "Memory Sparse Attention for Efficient End-to-End Memory Model Scaling to 100M Tokens" (Chen et al., 2025), integrated with predictive-coding energy minimization. Key components: 1. Router Projectors (W_QR, W_KR) for document-level relevance scoring 2. Chunk-wise KV compression via mean pooling 3. Top-k sparse document selection 4. Document-wise RoPE for extrapolation to 100M tokens 5. Memory Interleave for multi-hop reasoning via iterative settling 6. Contrastive auxiliary loss for router training The deep integration with PC-SHO-DLM: - Router scoring is part of the energy function (settling optimizes retrieval) - Precision heads and routers share information (uncertainty = retrieval need) - Unified mode: retrieval improves during settling as parameters update """ import math from dataclasses import dataclass, field from typing import Optional, Tuple, List import torch import torch.nn as nn import torch.nn.functional as F @dataclass class MSAConfig: """Configuration for Memory Sparse Attention.""" chunk_size: int = 64 # tokens per chunk for compression top_k: int = 16 # number of documents to retrieve router_dim: int = 128 # dimension of router projections n_router_heads: int = 8 # number of router heads apply_from_layer: int = 6 # only apply MSA to upper half of layers (MSA finding) aux_loss_weight: float = 0.1 # weight of contrastive routing loss aux_temperature: float = 0.05 # temperature for contrastive loss rope_base: float = 10000.0 # RoPE base frequency class DocumentWiseRoPE(nn.Module): """Document-wise Rotary Position Embedding. Each document gets independent position IDs starting from 0. Query tokens get global offset by top_k. This enables training on 64k but extrapolating to 100M tokens. """ def __init__(self, d_model: int, max_len: int = 8192, base: float = 10000.0): super().__init__() self.d_model = d_model self.max_len = max_len self.base = base # Precompute frequencies inv_freq = 1.0 / (base ** (torch.arange(0, d_model, 2).float() / d_model)) self.register_buffer("inv_freq", inv_freq) def _compute_rope(self, positions: torch.Tensor, dim: int) -> Tuple[torch.Tensor, torch.Tensor]: """Compute cos and sin for given position indices.""" # positions: (B, S) or (S,) freqs = torch.einsum("...s,d->...sd", positions.float(), self.inv_freq[:dim // 2].to(positions.device)) cos = freqs.cos() sin = freqs.sin() return cos, sin def apply_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """Apply rotary embeddings to input tensor.""" # x: (..., S, D), cos/sin: (..., S, D//2) d = x.shape[-1] x1, x2 = x[..., :d // 2], x[..., d // 2:] return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1) def forward(self, x: torch.Tensor, doc_boundaries: Optional[torch.Tensor] = None, global_offset: int = 0) -> torch.Tensor: """Apply document-wise RoPE. Args: x: (B, S, D) input doc_boundaries: (B, S) tensor of document IDs per position. Positions within the same doc get local IDs. If None, standard positional encoding. global_offset: offset for query positions (= top_k retrieved docs) """ B, S, D = x.shape if doc_boundaries is not None: # Document-wise: each doc starts at position 0 positions = torch.zeros(B, S, device=x.device, dtype=torch.long) for b in range(B): for doc_id in doc_boundaries[b].unique(): doc_mask = doc_boundaries[b] == doc_id positions[b, doc_mask] = torch.arange(doc_mask.sum(), device=x.device) else: # Standard positional encoding with optional offset positions = torch.arange(S, device=x.device).unsqueeze(0).expand(B, -1) + global_offset cos, sin = self._compute_rope(positions, D) return self.apply_rope(x, cos, sin) class RouterProjector(nn.Module): """Learned router for document relevance scoring. Separate from the main K/Q projectors — dedicated to retrieval. Produces routing keys (KR) and routing queries (QR) in a shared space. """ def __init__(self, d_model: int, router_dim: int, n_heads: int): super().__init__() self.n_heads = n_heads self.d_head = router_dim // n_heads self.q_proj = nn.Linear(d_model, router_dim) self.k_proj = nn.Linear(d_model, router_dim) def project_query(self, h_q: torch.Tensor) -> torch.Tensor: """Project query hidden states to routing space. (B, S, router_dim)""" return self.q_proj(h_q) def project_key(self, h_doc: torch.Tensor) -> torch.Tensor: """Project document hidden states to routing space. (B, S, router_dim)""" return self.k_proj(h_doc) class MemoryBank: """Stores compressed KV representations of documents. Offline: encode documents → chunk → mean-pool K,V,KR → store Online: query → router scores → top-k → load compressed KV The bank stores three things per document per layer: - K_bar: compressed keys (n_chunks, n_heads, d_head) - V_bar: compressed values (n_chunks, n_heads, d_head) - KR_bar: compressed routing keys (n_chunks, router_dim) """ def __init__(self, chunk_size: int = 64): self.chunk_size = chunk_size self.documents = {} # doc_id -> {layer_id -> {K_bar, V_bar, KR_bar}} self.doc_ids = [] def add_document(self, doc_id: str, layer_kvs: dict): """Add a document's compressed KV to the bank. Args: doc_id: unique identifier layer_kvs: {layer_idx: {"K": tensor, "V": tensor, "KR": tensor}} Each tensor has shape (n_chunks, ...) """ self.documents[doc_id] = layer_kvs if doc_id not in self.doc_ids: self.doc_ids.append(doc_id) def get_routing_keys(self, layer_idx: int) -> Tuple[torch.Tensor, List[str]]: """Get all routing keys for a layer. Returns (N_total_chunks, router_dim) + doc IDs.""" keys = [] ids = [] for doc_id in self.doc_ids: if layer_idx in self.documents[doc_id]: kr = self.documents[doc_id][layer_idx]["KR"] keys.append(kr) ids.extend([doc_id] * kr.shape[0]) if keys: return torch.cat(keys, dim=0), ids return None, [] def get_kv(self, doc_ids: List[str], layer_idx: int) -> Tuple[torch.Tensor, torch.Tensor]: """Get compressed K,V for selected documents at a layer.""" ks, vs = [], [] for did in doc_ids: if did in self.documents and layer_idx in self.documents[did]: ks.append(self.documents[did][layer_idx]["K"]) vs.append(self.documents[did][layer_idx]["V"]) if ks: return torch.cat(ks, dim=0), torch.cat(vs, dim=0) return None, None def __len__(self): return len(self.doc_ids) def chunk_mean_pool(x: torch.Tensor, chunk_size: int) -> torch.Tensor: """Compress sequence via chunk-wise mean pooling. Args: x: (B, S, D) or (S, D) chunk_size: tokens per chunk Returns: (B, n_chunks, D) or (n_chunks, D) """ if x.dim() == 2: S, D = x.shape n_chunks = math.ceil(S / chunk_size) # Pad to multiple of chunk_size if S % chunk_size != 0: pad = chunk_size - (S % chunk_size) x = F.pad(x, (0, 0, 0, pad)) return x.view(n_chunks, chunk_size, D).mean(dim=1) else: B, S, D = x.shape n_chunks = math.ceil(S / chunk_size) if S % chunk_size != 0: pad = chunk_size - (S % chunk_size) x = F.pad(x, (0, 0, 0, pad)) return x.view(B, n_chunks, chunk_size, D).mean(dim=2) class MSALayer(nn.Module): """Memory Sparse Attention layer. Replaces standard self-attention with sparse document-level retrieval. Applied only to upper layers (lower layers use standard attention). The key integration with PC-SHO-DLM: - Router scores become part of the energy function - Settling optimizes both hidden states AND retrieval quality - Precision heads inform router confidence """ def __init__(self, d_model: int, n_heads: int, d_ff: int, msa_config: MSAConfig, dropout: float = 0.1): super().__init__() self.d_model = d_model self.n_heads = n_heads self.d_head = d_model // n_heads self.msa_config = msa_config # Standard Q/K/V projectors self.q_proj = nn.Linear(d_model, d_model) self.k_proj = nn.Linear(d_model, d_model) self.v_proj = nn.Linear(d_model, d_model) self.o_proj = nn.Linear(d_model, d_model) # Router projectors (separate from main attention) self.router = RouterProjector(d_model, msa_config.router_dim, msa_config.n_router_heads) # Document-wise RoPE self.rope = DocumentWiseRoPE(self.d_head) # FFN self.ff = nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model), nn.Dropout(dropout), ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def compute_routing_scores(self, h_query: torch.Tensor, memory_routing_keys: torch.Tensor) -> torch.Tensor: """Score documents by relevance to query. Args: h_query: (B, S_q, D) query hidden states memory_routing_keys: (N_chunks, router_dim) compressed routing keys Returns: scores: (B, N_chunks) relevance scores """ # Project query to routing space qr = self.router.project_query(h_query) # (B, S_q, router_dim) # Normalize for cosine similarity qr_norm = F.normalize(qr, dim=-1) kr_norm = F.normalize(memory_routing_keys, dim=-1) # Score: max over query tokens, mean over heads # (B, S_q, router_dim) @ (N_chunks, router_dim)^T → (B, S_q, N_chunks) sim = torch.einsum("bsd,nd->bsn", qr_norm, kr_norm.to(qr_norm.device)) # Max-pool over query tokens scores = sim.max(dim=1).values # (B, N_chunks) return scores def sparse_attention(self, h_query: torch.Tensor, memory_k: torch.Tensor, memory_v: torch.Tensor, doc_boundaries: Optional[torch.Tensor] = None) -> torch.Tensor: """Attend to query + selected memory KV. Args: h_query: (B, S_q, D) query hidden states memory_k: (B, S_mem, D) compressed keys from selected documents memory_v: (B, S_mem, D) compressed values from selected documents Returns: output: (B, S_q, D) """ B, S_q, D = h_query.shape H, d = self.n_heads, self.d_head # Project query Q = self.q_proj(h_query).view(B, S_q, H, d).transpose(1, 2) if memory_k is not None: S_mem = memory_k.shape[1] # Concatenate memory + query KV K_mem = self.k_proj(memory_k).view(B, S_mem, H, d).transpose(1, 2) V_mem = self.v_proj(memory_v).view(B, S_mem, H, d).transpose(1, 2) K_q = self.k_proj(h_query).view(B, S_q, H, d).transpose(1, 2) V_q = self.v_proj(h_query).view(B, S_q, H, d).transpose(1, 2) # Apply document-wise RoPE to memory, global RoPE to query # (simplified: just offset query positions) K_ctx = torch.cat([K_mem, K_q], dim=2) # (B, H, S_mem+S_q, d) V_ctx = torch.cat([V_mem, V_q], dim=2) else: K_ctx = self.k_proj(h_query).view(B, S_q, H, d).transpose(1, 2) V_ctx = self.v_proj(h_query).view(B, S_q, H, d).transpose(1, 2) # Standard scaled dot-product attention scale = math.sqrt(d) attn = torch.einsum("bhsd,bhtd->bhst", Q, K_ctx) / scale attn = F.softmax(attn, dim=-1) attn = self.dropout(attn) out = torch.einsum("bhst,bhtd->bhsd", attn, V_ctx) out = out.transpose(1, 2).reshape(B, S_q, D) return self.o_proj(out) def forward(self, x: torch.Tensor, memory_k: Optional[torch.Tensor] = None, memory_v: Optional[torch.Tensor] = None) -> torch.Tensor: """Forward pass with optional memory context. If memory_k/v are provided, uses sparse attention over memory + local. Otherwise, falls back to standard bidirectional self-attention. """ residual = x x = self.norm1(x) x = residual + self.dropout(self.sparse_attention(x, memory_k, memory_v)) residual = x x = self.norm2(x) x = residual + self.ff(x) return x def compute_routing_aux_loss(scores_pos: torch.Tensor, scores_neg: torch.Tensor, temperature: float = 0.05) -> torch.Tensor: """Contrastive auxiliary loss for router training (MSA Eq. 5). Pushes positive document scores above negative document scores. Args: scores_pos: (B, n_pos) scores for relevant documents scores_neg: (B, n_neg) scores for irrelevant documents temperature: softmax temperature Returns: loss: scalar """ # For each positive, contrast against all negatives # L = -1/|P| * sum_i log(exp(s+_i/τ) / (exp(s+_i/τ) + sum_j exp(s-_j/τ))) pos_exp = (scores_pos / temperature).exp() # (B, n_pos) neg_exp_sum = (scores_neg / temperature).exp().sum(dim=-1, keepdim=True) # (B, 1) log_prob = (scores_pos / temperature) - torch.log(pos_exp + neg_exp_sum + 1e-10) loss = -log_prob.mean() return loss class MemoryEncoder: """Encodes documents into the memory bank. Offline process: runs each document through the model, extracts K, V, and KR at each MSA layer, compresses via chunk-wise mean pooling, and stores in the MemoryBank. """ @staticmethod @torch.no_grad() def encode_document(model, text: str, doc_id: str, memory_bank: MemoryBank, msa_layers: nn.ModuleList, chunk_size: int = 64, device: str = "cpu") -> None: """Encode a single document into the memory bank. Args: model: PCSHODLM model text: raw text of the document doc_id: unique identifier memory_bank: bank to store compressed KV msa_layers: list of MSALayer modules chunk_size: compression chunk size """ # Encode text as bytes tokens = torch.tensor( [min(b + 1, 256) for b in text.encode("utf-8")[:model.config.max_seq_len]], dtype=torch.long ).unsqueeze(0).to(device) # Pad if needed if tokens.shape[1] < model.config.max_seq_len: tokens = F.pad(tokens, (0, model.config.max_seq_len - tokens.shape[1])) # Run through model to get hidden states at each layer t = torch.ones(1, dtype=torch.long, device=device) # dummy timestep h = model.embed_input(tokens, t) layer_kvs = {} n_layers = len(model.forward_blocks) msa_start = n_layers // 2 # MSA applies to upper half for l, block in enumerate(model.forward_blocks): h = block(h) # For MSA layers, extract and compress KV + routing keys if l >= msa_start: msa_idx = l - msa_start if msa_idx >= len(msa_layers): continue msa_layer = msa_layers[msa_idx] # Project to K, V spaces K = msa_layer.k_proj(h.squeeze(0)) # (S, D) V = msa_layer.v_proj(h.squeeze(0)) # (S, D) KR = msa_layer.router.project_key(h.squeeze(0)) # (S, router_dim) # Compress via chunk-wise mean pooling K_bar = chunk_mean_pool(K, chunk_size) # (n_chunks, D) V_bar = chunk_mean_pool(V, chunk_size) KR_bar = chunk_mean_pool(KR, chunk_size) # (n_chunks, router_dim) layer_kvs[l] = {"K": K_bar, "V": V_bar, "KR": KR_bar} memory_bank.add_document(doc_id, layer_kvs) # ============================================================================= # Integration helper: create MSA layers for a PC-SHO-DLM model # ============================================================================= def create_msa_layers(model_config, msa_config: Optional[MSAConfig] = None) -> nn.ModuleList: """Create MSA layers for the upper half of the model. Args: model_config: PCSHOConfig msa_config: MSAConfig (uses defaults if None) Returns: ModuleList of MSALayer, one per upper-half layer """ if msa_config is None: msa_config = MSAConfig() n_msa_layers = model_config.n_layers - msa_config.apply_from_layer if n_msa_layers <= 0: n_msa_layers = model_config.n_layers // 2 layers = nn.ModuleList([ MSALayer( d_model=model_config.d_model, n_heads=model_config.n_heads, d_ff=model_config.d_ff, msa_config=msa_config, dropout=model_config.dropout, ) for _ in range(n_msa_layers) ]) return layers