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