Download src/msa.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 17.8 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/msa.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/src/msa.py
-
curl -L -o msa.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/msa.py
17.8 kB
| """ | |
| 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 | |
| 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. | |
| """ | |
| 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 | |