""" Direction F: Infinite Context via Recursive Settling The core insight: context window = settling budget, not sequence length. In standard transformers, context is bounded by the attention window. In PC-SHO-DLM + MSA, context is bounded by how many settling steps you can afford. Each settling round can retrieve new documents from the memory bank, so the effective context grows linearly with budget: |C_eff(K)| = min(K * B_doc, C_retain) where K = settling rounds, B_doc = docs retrieved per round, and C_retain = total documents in the memory bank. The recursive settling loop: 1. Settle on current context (query + retrieved docs) 2. Check energy convergence: |E^{k+1} - E^k| < epsilon 3. If not converged, use updated hidden states to re-route and retrieve additional documents 4. Repeat until converged or budget exhausted This is a natural consequence of the PC architecture: settling IS inference, and each settling step can refine what the model attends to. """ import math from dataclasses import dataclass, field from typing import Optional, Tuple, List, Dict import torch import torch.nn as nn import torch.nn.functional as F from model import PCSHODLM, PCSHOConfig from msa import ( MSALayer, MSAConfig, MemoryBank, MemoryEncoder, RouterProjector, chunk_mean_pool, create_msa_layers, ) @dataclass class InfiniteContextResult: """Result of infinite-context processing.""" logits: torch.Tensor # (B, S, V) final output logits effective_context_size: int # total document chunks accessed documents_accessed: List[str] # ordered list of document IDs retrieved energy_trace: List[float] # energy at each settling round rounds_used: int # settling rounds before convergence converged: bool # whether energy converged within budget per_round_retrievals: List[int] # new docs retrieved per round class InfiniteContextProcessor: """Infinite context via recursive settling. Context window = settling budget, not sequence length. Given a query and a massive memory bank (1000+ documents), this processor iteratively: 1. Settles on the current context 2. Uses the settled hidden states to re-route into the memory bank 3. Retrieves new relevant documents not yet seen 4. Settles again with the expanded context Convergence detection: stop when |E^{k+1} - E^k| < epsilon. The bound on effective context: |C_eff(K)| = min(K * B_doc, C_retain) where K = settling rounds used, B_doc = documents per retrieval, C_retain = total available documents. Args: model: PC-SHO-DLM model instance msa_layers: MSA layers (upper-half attention with routing) memory_bank: pre-encoded document bank docs_per_round: number of new documents to retrieve each round epsilon: energy convergence threshold max_budget: maximum total settling rounds inner_settling_steps: PC settling iterations per round """ def __init__( self, model: PCSHODLM, msa_layers: nn.ModuleList, memory_bank: MemoryBank, docs_per_round: int = 4, epsilon: float = 1e-3, max_budget: int = 50, inner_settling_steps: int = 4, ): self.model = model self.msa_layers = msa_layers self.memory_bank = memory_bank self.docs_per_round = docs_per_round self.epsilon = epsilon self.max_budget = max_budget self.inner_settling_steps = inner_settling_steps def _tokenize_query(self, query: str, device: torch.device) -> torch.Tensor: """Byte-level tokenization matching MemoryEncoder conventions.""" tokens = torch.tensor( [min(b + 1, 256) for b in query.encode("utf-8")[:self.model.config.max_seq_len]], dtype=torch.long, ).unsqueeze(0).to(device) if tokens.shape[1] < self.model.config.max_seq_len: tokens = F.pad(tokens, (0, self.model.config.max_seq_len - tokens.shape[1])) return tokens def _retrieve_documents( self, h_query: torch.Tensor, already_retrieved: set, layer_idx: int, n_docs: int, ) -> List[str]: """Route into memory bank and retrieve top-n unseen documents. Uses the MSA router to score all document chunks, then selects the top-scoring documents not already in the active context. Args: h_query: (B, S, D) current hidden states at the routing layer already_retrieved: set of doc IDs already retrieved layer_idx: which model layer to use for routing n_docs: how many new documents to retrieve Returns: List of newly retrieved document IDs """ if len(self.memory_bank) == 0: return [] # Get routing keys from the memory bank for this layer routing_keys, chunk_doc_ids = self.memory_bank.get_routing_keys(layer_idx) if routing_keys is None or len(chunk_doc_ids) == 0: return [] # Find the MSA layer for routing n_layers = len(self.model.forward_blocks) msa_start = n_layers // 2 msa_idx = layer_idx - msa_start if msa_idx < 0 or msa_idx >= len(self.msa_layers): return [] msa_layer = self.msa_layers[msa_idx] # Score all chunks with torch.no_grad(): scores = msa_layer.compute_routing_scores( h_query, routing_keys.to(h_query.device) ) # (B, N_chunks) # Aggregate chunk scores to document scores doc_scores: Dict[str, float] = {} scores_flat = scores[0].cpu().tolist() # batch dim = 0 for i, doc_id in enumerate(chunk_doc_ids): if doc_id in already_retrieved: continue if doc_id not in doc_scores: doc_scores[doc_id] = 0.0 doc_scores[doc_id] = max(doc_scores[doc_id], scores_flat[i]) if not doc_scores: return [] # Sort by score descending, take top n ranked = sorted(doc_scores.items(), key=lambda x: x[1], reverse=True) return [doc_id for doc_id, _ in ranked[:n_docs]] def _gather_memory_kv( self, doc_ids: List[str], layer_idx: int, device: torch.device, ) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: """Gather compressed K, V tensors for the retrieved documents. Returns: memory_k: (1, total_chunks, D) or None memory_v: (1, total_chunks, D) or None """ K, V = self.memory_bank.get_kv(doc_ids, layer_idx) if K is None: return None, None # Add batch dimension return K.unsqueeze(0).to(device), V.unsqueeze(0).to(device) def _run_settling_round( self, tokens: torch.Tensor, retrieved_docs: List[str], prev_h: Optional[List[torch.Tensor]], prev_v: Optional[List[torch.Tensor]], ) -> Tuple[List[torch.Tensor], List[torch.Tensor], float]: """Run one round of PC settling with the current retrieved context. This performs the inner settling loop (multiple PC iterations) using the MSA layers to attend over the retrieved documents. Returns: h_settled: settled hidden states v_final: final velocities final_energy: energy after settling """ model = self.model config = model.config device = tokens.device B, S = tokens.shape L = config.n_layers n_layers = len(model.forward_blocks) msa_start = n_layers // 2 # Timestep (use t=1 for inference-like settling) t = torch.ones(B, dtype=torch.long, device=device) # Create a mask (no tokens masked -- pure inference) mask = torch.zeros(B, S, dtype=torch.bool, device=device) # Embed input h_0 = model.embed_input(tokens, t) # Forward pass with MSA integration: # Lower layers use standard forward blocks, # upper layers use MSA with retrieved document context. h = h_0 h_init = [h_0] for l, block in enumerate(model.forward_blocks): if l >= msa_start and retrieved_docs: msa_idx = l - msa_start if msa_idx < len(self.msa_layers): # Get memory KV for this layer mem_k, mem_v = self._gather_memory_kv( retrieved_docs, l, device ) # Run MSA layer with memory context h = self.msa_layers[msa_idx](h, memory_k=mem_k, memory_v=mem_v) else: h = block(h) else: h = block(h) h_init.append(h) # Now run PC settling on these initialized states # Override settling steps for this inner loop orig_steps = config.n_settling_steps config.n_settling_steps = self.inner_settling_steps h_settled, v_final, energies, _, _ = model.settle( h_init, tokens, mask, t, prev_h=prev_h, prev_v=prev_v, ) config.n_settling_steps = orig_steps final_energy = energies[-1] if energies else float("inf") return h_settled, v_final, final_energy def process( self, query: str, device: str = "cpu", ) -> InfiniteContextResult: """Process a query against the full memory bank via recursive settling. The main loop: for round in range(max_budget): settle(query + retrieved_docs) if converged: break retrieve_more_docs(using settled hidden states) Args: query: input text device: torch device Returns: InfiniteContextResult with logits, documents accessed, energy trace, etc. """ model = self.model device = torch.device(device) model = model.to(device) for layer in self.msa_layers: layer.to(device) tokens = self._tokenize_query(query, device) n_layers = len(model.forward_blocks) msa_start = n_layers // 2 # Use the first MSA-eligible layer for routing routing_layer = msa_start retrieved_docs: List[str] = [] retrieved_set: set = set() energy_trace: List[float] = [] per_round_retrievals: List[int] = [] prev_h = None prev_v = None converged = False with torch.no_grad(): for round_idx in range(self.max_budget): # --- Retrieve new documents --- if round_idx == 0: # First round: use amortized forward pass for initial routing t_dummy = torch.ones(1, dtype=torch.long, device=device) h_0 = model.embed_input(tokens, t_dummy) h_for_routing = h_0 for l in range(routing_layer + 1): h_for_routing = model.forward_blocks[l](h_for_routing) else: # Subsequent rounds: use settled states at the routing layer h_for_routing = prev_h[routing_layer + 1] if prev_h else None if h_for_routing is not None: new_docs = self._retrieve_documents( h_for_routing, retrieved_set, routing_layer, self.docs_per_round, ) retrieved_docs.extend(new_docs) retrieved_set.update(new_docs) per_round_retrievals.append(len(new_docs)) else: per_round_retrievals.append(0) # --- Settle with current context --- h_settled, v_final, energy = self._run_settling_round( tokens, retrieved_docs, prev_h, prev_v, ) energy_trace.append(energy) prev_h = h_settled prev_v = v_final # --- Check convergence --- if len(energy_trace) >= 2: delta = abs(energy_trace[-1] - energy_trace[-2]) if delta < self.epsilon: converged = True break # --- Check if memory bank exhausted --- if len(retrieved_set) >= len(self.memory_bank): break # Final readout with torch.no_grad(): logits = model.readout(model.readout_norm(prev_h[-1])) return InfiniteContextResult( logits=logits, effective_context_size=len(retrieved_docs), documents_accessed=list(retrieved_docs), energy_trace=energy_trace, rounds_used=len(energy_trace), converged=converged, per_round_retrievals=per_round_retrievals, ) def stream_process( self, text_chunks: List[str], device: str = "cpu", ) -> List[InfiniteContextResult]: """Process an arbitrarily long text stream, maintaining state via continuation. Each chunk is processed using the recursive settling loop, with hidden states and velocities carried forward from the previous chunk. This enables processing text of unlimited length without ever exceeding the model's sequence window. The key mechanism: the recurrent state (h, v) from settling compresses all prior context into a fixed-size representation. New chunks are settled starting from this warm-start state, so information from earlier chunks persists through the second-order dynamics. Args: text_chunks: list of text segments to process in order device: torch device Returns: List of InfiniteContextResult, one per chunk """ model = self.model device = torch.device(device) model = model.to(device) for layer in self.msa_layers: layer.to(device) results: List[InfiniteContextResult] = [] continuation_h: Optional[List[torch.Tensor]] = None continuation_v: Optional[List[torch.Tensor]] = None # Accumulate which documents have been accessed across chunks cumulative_docs: List[str] = [] cumulative_set: set = set() for chunk_idx, chunk_text in enumerate(text_chunks): tokens = self._tokenize_query(chunk_text, device) n_layers = len(model.forward_blocks) msa_start = n_layers // 2 routing_layer = msa_start # Per-chunk retrieval state chunk_docs: List[str] = [] energy_trace: List[float] = [] per_round_retrievals: List[int] = [] converged = False prev_h = continuation_h prev_v = continuation_v with torch.no_grad(): for round_idx in range(self.max_budget): # Route using current state if prev_h is not None and round_idx > 0: h_for_routing = prev_h[routing_layer + 1] else: t_dummy = torch.ones(1, dtype=torch.long, device=device) h_0 = model.embed_input(tokens, t_dummy) h_for_routing = h_0 for l in range(routing_layer + 1): h_for_routing = model.forward_blocks[l](h_for_routing) new_docs = self._retrieve_documents( h_for_routing, cumulative_set, routing_layer, self.docs_per_round, ) chunk_docs.extend(new_docs) cumulative_docs.extend(new_docs) cumulative_set.update(new_docs) per_round_retrievals.append(len(new_docs)) # Settle h_settled, v_final, energy = self._run_settling_round( tokens, cumulative_docs, # use ALL docs seen so far across stream prev_h, prev_v, ) energy_trace.append(energy) prev_h = h_settled prev_v = v_final # Convergence check if len(energy_trace) >= 2: delta = abs(energy_trace[-1] - energy_trace[-2]) if delta < self.epsilon: converged = True break if len(cumulative_set) >= len(self.memory_bank): break # Carry state forward to next chunk continuation_h = prev_h continuation_v = [ self.model.config.velocity_decay * v for v in prev_v ] if prev_v else None # Readout for this chunk with torch.no_grad(): logits = model.readout(model.readout_norm(prev_h[-1])) results.append(InfiniteContextResult( logits=logits, effective_context_size=len(cumulative_docs), documents_accessed=list(chunk_docs), energy_trace=energy_trace, rounds_used=len(energy_trace), converged=converged, per_round_retrievals=per_round_retrievals, )) return results @staticmethod def effective_context_bound( settling_rounds: int, docs_per_round: int, total_docs: int, ) -> int: """Compute the theoretical bound on effective context. |C_eff(K)| = min(K * B_doc, C_retain) Args: settling_rounds: K, the number of settling rounds docs_per_round: B_doc, documents retrieved per round total_docs: C_retain, total documents available Returns: Upper bound on effective context size """ return min(settling_rounds * docs_per_round, total_docs) # ============================================================================ # Test: effective context growth with settling budget # ============================================================================ def test_infinite_context(): """Demonstrate that effective context grows with settling budget. Creates a model with 60 documents in the memory bank and shows how the number of accessed documents increases as we allow more settling rounds. """ print("=" * 72) print("Direction F: Infinite Context via Recursive Settling") print("=" * 72) # -- Setup: small model for testing -- config = PCSHOConfig( vocab_size=300, max_seq_len=64, d_model=128, n_heads=4, n_layers=4, d_ff=256, n_diffusion_steps=100, n_settling_steps=4, feedback_rank=32, ) model = PCSHODLM(config) model.eval() msa_config = MSAConfig( chunk_size=16, top_k=4, router_dim=64, n_router_heads=4, ) msa_layers = create_msa_layers(config, msa_config) # -- Build a memory bank with 60 documents -- n_docs = 60 memory_bank = MemoryBank(chunk_size=msa_config.chunk_size) print(f"\nEncoding {n_docs} documents into memory bank...") for i in range(n_docs): # Create synthetic documents with distinct content doc_text = f"Document {i}: " + f"topic-{i % 10} " * 20 MemoryEncoder.encode_document( model, doc_text, f"doc_{i:04d}", memory_bank, msa_layers, chunk_size=msa_config.chunk_size, ) print(f"Memory bank size: {len(memory_bank)} documents") # -- Test: vary settling budget and observe context growth -- query = "Find information about topic-3 and topic-7" budgets = [1, 2, 5, 10, 20, 30] docs_per_round = 4 print(f"\nQuery: \"{query}\"") print(f"Docs per round: {docs_per_round}") print(f"\n{'Budget':>8} {'Rounds':>8} {'Docs Accessed':>15} " f"{'Converged':>10} {'Bound':>8} {'Final Energy':>14}") print("-" * 72) for budget in budgets: processor = InfiniteContextProcessor( model=model, msa_layers=msa_layers, memory_bank=memory_bank, docs_per_round=docs_per_round, epsilon=1e-4, max_budget=budget, inner_settling_steps=2, ) result = processor.process(query, device="cpu") bound = InfiniteContextProcessor.effective_context_bound( budget, docs_per_round, n_docs, ) final_e = result.energy_trace[-1] if result.energy_trace else float("nan") print( f"{budget:>8} {result.rounds_used:>8} " f"{result.effective_context_size:>15} " f"{'yes' if result.converged else 'no':>10} " f"{bound:>8} {final_e:>14.2f}" ) # -- Test: stream processing of long text -- print("\n" + "=" * 72) print("Stream Processing: arbitrarily long input via continuation") print("=" * 72) chunks = [ "What is the relationship between topic-3 and topic-7?", "Also consider how topic-1 and topic-5 interact with them.", "Finally, summarize the connections across all topics.", ] processor = InfiniteContextProcessor( model=model, msa_layers=msa_layers, memory_bank=memory_bank, docs_per_round=3, epsilon=1e-4, max_budget=8, inner_settling_steps=2, ) print(f"\nProcessing {len(chunks)} text chunks in streaming mode...") results = processor.stream_process(chunks, device="cpu") for i, (chunk, result) in enumerate(zip(chunks, results)): print(f"\n Chunk {i + 1}: \"{chunk[:50]}...\"") print(f" Rounds: {result.rounds_used}, " f"Docs this chunk: {len(result.documents_accessed)}, " f"Cumulative context: {result.effective_context_size}, " f"Converged: {result.converged}") if result.energy_trace: print(f" Energy trace: [{', '.join(f'{e:.2f}' for e in result.energy_trace)}]") # -- Verify the bound holds -- print("\n" + "=" * 72) print("Bound Verification: |C_eff(K)| = min(K * B_doc, C_retain)") print("=" * 72) all_passed = True for budget in budgets: processor = InfiniteContextProcessor( model=model, msa_layers=msa_layers, memory_bank=memory_bank, docs_per_round=docs_per_round, epsilon=1e-4, max_budget=budget, inner_settling_steps=2, ) result = processor.process(query, device="cpu") bound = InfiniteContextProcessor.effective_context_bound( result.rounds_used, docs_per_round, n_docs, ) holds = result.effective_context_size <= bound all_passed = all_passed and holds status = "PASS" if holds else "FAIL" print(f" Budget={budget:>3}: accessed={result.effective_context_size:>3}, " f"bound={bound:>3} [{status}]") print(f"\nAll bounds hold: {all_passed}") print("\nDone.") if __name__ == "__main__": test_infinite_context()