pc-sho-dlm-code / src /msa.py
zotowata's picture
Upload src/msa.py with huggingface_hub
24af195 unverified
Raw History Blame Contribute Delete
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
@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