pc-sho-dlm-code / src /infinite_context.py
Zae
ric
49589ef
Raw History Blame Contribute Delete
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,
)
@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()