Download src/self_improving.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 37.1 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/self_improving.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/src/self_improving.py
-
curl -L -o self_improving.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/self_improving.py
37.1 kB
| """ | |
| Direction G: Self-Improving Retrieval for PC-SHO-DLM + MSA | |
| A retrieval system that improves with every query -- no explicit retraining. | |
| Each query now uses the repaired unified path: | |
| 1. Settle hidden states toward a low-energy solution (fast timescale) | |
| 2. Apply model updates once from the settled state (slow timescale) | |
| 3. Update router parameters from the settled retrieval signal | |
| Convergence guarantee (Borkar 2008 two-timescale + PC contraction): | |
| E[||W_QR^N - W_QR*||^2] = O(1 / sqrt(N)) | |
| Key insight: predictive-coding settling supplies the local error signals, | |
| but the stable training rule is to update model parameters from settled | |
| states rather than from transient microsteps. Router weights still adapt | |
| online from the retrieval signal within the query. | |
| Safety mechanisms: | |
| - Elastic regularization: prevents catastrophic drift from initial weights | |
| - Snapshot/rollback: revert if quality degrades | |
| - Drift monitoring: ||theta_n - theta_0|| / ||theta_0|| tracked per query | |
| """ | |
| import copy | |
| 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, InferenceUpdater | |
| from msa import ( | |
| MSAConfig, MSALayer, MemoryBank, MemoryEncoder, | |
| RouterProjector, create_msa_layers, chunk_mean_pool, | |
| compute_routing_aux_loss, | |
| ) | |
| # ============================================================================ | |
| # Configuration | |
| # ============================================================================ | |
| class SelfImprovingConfig: | |
| """Configuration for the self-improving retrieval system.""" | |
| # Elastic regularization | |
| elastic_lambda: float = 0.01 # L_elastic = lambda * ||theta - theta_0||^2 | |
| drift_threshold: float = 0.10 # activate elastic reg when drift > 10% | |
| drift_hard_cap: float = 0.30 # force rollback if drift > 30% | |
| # Unified settling for retrieval queries | |
| n_settling_steps: int = 6 # inner loop iterations per query | |
| param_lr_scale: float = 0.01 # slow timescale for parameters | |
| # Quality tracking | |
| quality_ema_alpha: float = 0.1 # exponential moving average smoothing | |
| quality_window: int = 10 # window for rolling average | |
| # Snapshot policy | |
| snapshot_every: int = 10 # save snapshot every N queries | |
| max_snapshots: int = 5 # keep at most this many snapshots | |
| # Router-specific learning rate scaling | |
| router_lr_boost: float = 2.0 # router params get boosted LR | |
| readout_lr_scale: float = 0.5 # readout params get reduced LR | |
| # ============================================================================ | |
| # Retrieval Quality Metric | |
| # ============================================================================ | |
| class RetrievalQualityTracker: | |
| """Tracks Q_N = retrieval quality at query N. | |
| Quality is measured as a composite of: | |
| - Router confidence: max routing score for selected documents | |
| - Settling energy reduction: E_final / E_initial (lower is better) | |
| - Answer coherence: negative entropy of output distribution | |
| All three are normalized to [0, 1] and combined. | |
| """ | |
| def __init__(self, ema_alpha: float = 0.1): | |
| self.ema_alpha = ema_alpha | |
| self.history: List[float] = [] | |
| self.components: List[Dict[str, float]] = [] | |
| self._ema = 0.0 | |
| self._initialized = False | |
| def record(self, router_confidence: float, energy_ratio: float, | |
| answer_coherence: float) -> float: | |
| """Record quality for one query and return composite Q_N.""" | |
| # Router confidence: already in [0, 1] (cosine similarity based) | |
| q_router = max(0.0, min(1.0, router_confidence)) | |
| # Energy ratio: E_final / E_initial. Lower = better settling. | |
| # Map to [0, 1] where 1 = perfect settling (ratio -> 0) | |
| q_energy = max(0.0, min(1.0, 1.0 - energy_ratio)) | |
| # Answer coherence: negative entropy normalized by log(vocab_size) | |
| # Higher coherence (lower entropy) = better. Already in [0, 1]. | |
| q_coherence = max(0.0, min(1.0, answer_coherence)) | |
| # Composite: weighted average | |
| q_n = 0.4 * q_router + 0.3 * q_energy + 0.3 * q_coherence | |
| self.history.append(q_n) | |
| self.components.append({ | |
| "router_confidence": q_router, | |
| "energy_reduction": q_energy, | |
| "answer_coherence": q_coherence, | |
| "composite": q_n, | |
| }) | |
| # Update EMA | |
| if not self._initialized: | |
| self._ema = q_n | |
| self._initialized = True | |
| else: | |
| self._ema = self.ema_alpha * q_n + (1 - self.ema_alpha) * self._ema | |
| return q_n | |
| def current_quality(self) -> float: | |
| return self._ema if self._initialized else 0.0 | |
| def n_queries(self) -> int: | |
| return len(self.history) | |
| def get_improvement_curve(self) -> List[float]: | |
| """Return the full Q_N sequence.""" | |
| return list(self.history) | |
| def get_rolling_average(self, window: int = 10) -> List[float]: | |
| """Return rolling average of quality for smoother visualization.""" | |
| if len(self.history) < window: | |
| return list(self.history) | |
| result = [] | |
| for i in range(len(self.history)): | |
| start = max(0, i - window + 1) | |
| result.append(sum(self.history[start:i + 1]) / (i - start + 1)) | |
| return result | |
| # ============================================================================ | |
| # Drift Monitor | |
| # ============================================================================ | |
| class DriftMonitor: | |
| """Monitors parameter drift: drift_n = ||theta_n - theta_0|| / ||theta_0||. | |
| Provides per-component drift (router, forward blocks, feedback, readout) | |
| and aggregate drift for the elastic regularization trigger. | |
| """ | |
| def __init__(self): | |
| self._theta_0: Optional[Dict[str, torch.Tensor]] = None | |
| self._theta_0_norm: float = 0.0 | |
| self.history: List[float] = [] | |
| self.component_history: List[Dict[str, float]] = [] | |
| def set_baseline(self, model: nn.Module, msa_layers: nn.ModuleList) -> None: | |
| """Snapshot initial parameters as theta_0.""" | |
| self._theta_0 = {} | |
| total_norm_sq = 0.0 | |
| for name, p in model.named_parameters(): | |
| self._theta_0[f"model.{name}"] = p.data.clone() | |
| total_norm_sq += p.data.norm().item() ** 2 | |
| for name, p in msa_layers.named_parameters(): | |
| self._theta_0[f"msa.{name}"] = p.data.clone() | |
| total_norm_sq += p.data.norm().item() ** 2 | |
| self._theta_0_norm = math.sqrt(total_norm_sq) | |
| def compute_drift(self, model: nn.Module, msa_layers: nn.ModuleList) -> float: | |
| """Compute current drift from baseline. Returns scalar drift ratio.""" | |
| if self._theta_0 is None: | |
| return 0.0 | |
| delta_sq = 0.0 | |
| component_deltas = {"router": 0.0, "forward": 0.0, "feedback": 0.0, "other": 0.0} | |
| for name, p in model.named_parameters(): | |
| key = f"model.{name}" | |
| if key in self._theta_0: | |
| d = (p.data - self._theta_0[key].to(p.device)).norm().item() ** 2 | |
| delta_sq += d | |
| if "forward_blocks" in name: | |
| component_deltas["forward"] += d | |
| elif "feedback_blocks" in name: | |
| component_deltas["feedback"] += d | |
| else: | |
| component_deltas["other"] += d | |
| for name, p in msa_layers.named_parameters(): | |
| key = f"msa.{name}" | |
| if key in self._theta_0: | |
| d = (p.data - self._theta_0[key].to(p.device)).norm().item() ** 2 | |
| delta_sq += d | |
| if "router" in name: | |
| component_deltas["router"] += d | |
| else: | |
| component_deltas["other"] += d | |
| drift = math.sqrt(delta_sq) / max(self._theta_0_norm, 1e-10) | |
| self.history.append(drift) | |
| # Normalize component deltas | |
| component_drift = { | |
| k: math.sqrt(v) / max(self._theta_0_norm, 1e-10) | |
| for k, v in component_deltas.items() | |
| } | |
| self.component_history.append(component_drift) | |
| return drift | |
| def get_elastic_penalty(self, model: nn.Module, msa_layers: nn.ModuleList, | |
| lam: float) -> torch.Tensor: | |
| """Compute L_elastic = lambda * ||theta - theta_0||^2. | |
| Returns a differentiable scalar loss to be added to the energy. | |
| """ | |
| if self._theta_0 is None: | |
| return torch.tensor(0.0) | |
| penalty = torch.tensor(0.0, device=next(model.parameters()).device) | |
| for name, p in model.named_parameters(): | |
| key = f"model.{name}" | |
| if key in self._theta_0: | |
| penalty = penalty + (p - self._theta_0[key].to(p.device)).pow(2).sum() | |
| for name, p in msa_layers.named_parameters(): | |
| key = f"msa.{name}" | |
| if key in self._theta_0: | |
| penalty = penalty + (p - self._theta_0[key].to(p.device)).pow(2).sum() | |
| return lam * penalty | |
| # ============================================================================ | |
| # Self-Improving Retriever | |
| # ============================================================================ | |
| class SelfImprovingRetriever: | |
| """A retrieval system that improves with every query. | |
| Wraps PC-SHO-DLM model + MSA layers + MemoryBank into a unified | |
| retrieval engine where each query triggers settling, then applies | |
| post-settle model updates and router adaptation. | |
| Convergence bound: | |
| E[||W_QR^N - W_QR*||^2] = O(1 / sqrt(N)) | |
| This follows from Borkar (2008) two-timescale stochastic approximation: | |
| the fast process (hidden state settling) converges at rate O(1/K) per query, | |
| while the slow process (parameter updates) converges at rate O(1/sqrt(N)) | |
| over queries, because the effective noise variance is bounded by the | |
| settling residual which contracts geometrically. | |
| Usage: | |
| retriever = SelfImprovingRetriever(model, msa_layers, memory_bank, config) | |
| for text in queries: | |
| answer = retriever.query(text) | |
| curve = retriever.get_improvement_curve() | |
| """ | |
| def __init__( | |
| self, | |
| model: PCSHODLM, | |
| msa_layers: nn.ModuleList, | |
| memory_bank: MemoryBank, | |
| config: Optional[SelfImprovingConfig] = None, | |
| device: str = "cpu", | |
| ): | |
| self.model = model | |
| self.msa_layers = msa_layers | |
| self.memory_bank = memory_bank | |
| self.config = config or SelfImprovingConfig() | |
| self.device = device | |
| # Core tracking | |
| self.quality_tracker = RetrievalQualityTracker( | |
| ema_alpha=self.config.quality_ema_alpha | |
| ) | |
| self.drift_monitor = DriftMonitor() | |
| self.drift_monitor.set_baseline(model, msa_layers) | |
| # Query counter | |
| self._query_count = 0 | |
| # Snapshot management | |
| self._snapshots: List[Dict[str, torch.Tensor]] = [] | |
| self._snapshot_queries: List[int] = [] | |
| self._save_snapshot() # initial snapshot | |
| # Energy history per query (for diagnostics) | |
| self.energy_traces: List[List[float]] = [] | |
| # ------------------------------------------------------------------ | |
| # Snapshot / Rollback | |
| # ------------------------------------------------------------------ | |
| def _save_snapshot(self) -> None: | |
| """Save current parameters as a snapshot.""" | |
| snapshot = {} | |
| for name, p in self.model.named_parameters(): | |
| snapshot[f"model.{name}"] = p.data.clone() | |
| for name, p in self.msa_layers.named_parameters(): | |
| snapshot[f"msa.{name}"] = p.data.clone() | |
| self._snapshots.append(snapshot) | |
| self._snapshot_queries.append(self._query_count) | |
| # Prune old snapshots | |
| while len(self._snapshots) > self.config.max_snapshots: | |
| self._snapshots.pop(0) | |
| self._snapshot_queries.pop(0) | |
| def _restore_snapshot(self, idx: int = -1) -> None: | |
| """Restore parameters from a snapshot.""" | |
| snapshot = self._snapshots[idx] | |
| for name, p in self.model.named_parameters(): | |
| key = f"model.{name}" | |
| if key in snapshot: | |
| p.data.copy_(snapshot[key]) | |
| for name, p in self.msa_layers.named_parameters(): | |
| key = f"msa.{name}" | |
| if key in snapshot: | |
| p.data.copy_(snapshot[key]) | |
| def reset_to_original(self) -> None: | |
| """Reset all parameters to the original (query 0) state.""" | |
| self._restore_snapshot(0) | |
| self._query_count = 0 | |
| self.quality_tracker = RetrievalQualityTracker( | |
| ema_alpha=self.config.quality_ema_alpha | |
| ) | |
| self.drift_monitor = DriftMonitor() | |
| self.drift_monitor.set_baseline(self.model, self.msa_layers) | |
| self._snapshots = self._snapshots[:1] | |
| self._snapshot_queries = self._snapshot_queries[:1] | |
| self.energy_traces = [] | |
| def save_state(self, path: str) -> None: | |
| """Save full retriever state to disk.""" | |
| state = { | |
| "model_state": self.model.state_dict(), | |
| "msa_state": self.msa_layers.state_dict(), | |
| "quality_history": self.quality_tracker.history, | |
| "quality_components": self.quality_tracker.components, | |
| "drift_history": self.drift_monitor.history, | |
| "drift_components": self.drift_monitor.component_history, | |
| "query_count": self._query_count, | |
| "energy_traces": self.energy_traces, | |
| "config": self.config, | |
| } | |
| torch.save(state, path) | |
| def load_state(self, path: str) -> None: | |
| """Load retriever state from disk.""" | |
| state = torch.load(path, map_location=self.device, weights_only=False) | |
| self.model.load_state_dict(state["model_state"]) | |
| self.msa_layers.load_state_dict(state["msa_state"]) | |
| self.quality_tracker.history = state["quality_history"] | |
| self.quality_tracker.components = state["quality_components"] | |
| self.drift_monitor.history = state["drift_history"] | |
| self.drift_monitor.component_history = state["drift_components"] | |
| self._query_count = state["query_count"] | |
| self.energy_traces = state["energy_traces"] | |
| if "config" in state: | |
| self.config = state["config"] | |
| # ------------------------------------------------------------------ | |
| # Core: Query Processing with Self-Improvement | |
| # ------------------------------------------------------------------ | |
| def _tokenize(self, text: str) -> torch.Tensor: | |
| """Simple byte-level tokenization (matches MemoryEncoder).""" | |
| tokens = [min(b + 1, self.model.config.vocab_size - 1) | |
| for b in text.encode("utf-8")[:self.model.config.max_seq_len]] | |
| t = torch.tensor(tokens, dtype=torch.long, device=self.device).unsqueeze(0) | |
| if t.shape[1] < self.model.config.max_seq_len: | |
| t = F.pad(t, (0, self.model.config.max_seq_len - t.shape[1])) | |
| return t | |
| def _retrieve_documents(self, h_query: torch.Tensor, layer_idx: int | |
| ) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor], | |
| float, List[str]]: | |
| """Route query through MSA to retrieve relevant documents. | |
| Returns: | |
| memory_k: compressed keys from top-k docs (or None) | |
| memory_v: compressed values from top-k docs (or None) | |
| router_confidence: max routing score (for quality tracking) | |
| selected_ids: IDs of selected documents | |
| """ | |
| if len(self.memory_bank) == 0: | |
| return None, None, 0.0, [] | |
| msa_start = len(self.model.forward_blocks) // 2 | |
| msa_idx = layer_idx - msa_start | |
| if msa_idx < 0 or msa_idx >= len(self.msa_layers): | |
| return None, None, 0.0, [] | |
| msa_layer = self.msa_layers[msa_idx] | |
| # Get all routing keys from memory bank | |
| routing_keys, chunk_doc_ids = self.memory_bank.get_routing_keys(layer_idx) | |
| if routing_keys is None: | |
| return None, None, 0.0, [] | |
| routing_keys = routing_keys.to(self.device) | |
| # Compute routing scores | |
| scores = msa_layer.compute_routing_scores(h_query, routing_keys) # (1, N_chunks) | |
| # Top-k selection | |
| k = min(self.msa_layers[0].msa_config.top_k, scores.shape[1]) | |
| top_scores, top_indices = scores.topk(k, dim=1) | |
| router_confidence = top_scores.max().item() | |
| # Map chunk indices to document IDs | |
| selected_doc_ids = list(set( | |
| chunk_doc_ids[idx.item()] for idx in top_indices[0] | |
| if idx.item() < len(chunk_doc_ids) | |
| )) | |
| # Load compressed KV for selected documents | |
| memory_k, memory_v = self.memory_bank.get_kv(selected_doc_ids, layer_idx) | |
| if memory_k is not None: | |
| memory_k = memory_k.unsqueeze(0).to(self.device) # (1, S_mem, D) | |
| memory_v = memory_v.unsqueeze(0).to(self.device) | |
| return memory_k, memory_v, router_confidence, selected_doc_ids | |
| def query(self, text: str) -> Dict: | |
| """Process a query with self-improving retrieval. | |
| This is the main entry point. Each call: | |
| 1. Tokenizes the query | |
| 2. Runs the forward pass through lower layers | |
| 3. At MSA layers, routes to memory bank and retrieves documents | |
| 4. Runs shared settling, then post-settle model and router updates | |
| 5. Decodes the answer | |
| 6. Tracks quality, drift, and applies elastic regularization if needed | |
| Args: | |
| text: query text | |
| Returns: | |
| dict with keys: answer_logits, answer_tokens, quality, drift, | |
| energy_trace, retrieved_docs | |
| """ | |
| self._query_count += 1 | |
| config = self.config | |
| model = self.model | |
| mc = model.config | |
| # Tokenize | |
| tokens = self._tokenize(text) | |
| B, S = tokens.shape | |
| # Create a partial mask: treat last 25% of non-padding tokens as | |
| # "to predict" (simulates the query -> answer pattern) | |
| non_pad = (tokens != 0).sum(dim=1).item() | |
| mask_start = max(1, int(non_pad * 0.75)) | |
| mask = torch.zeros(B, S, dtype=torch.bool, device=self.device) | |
| mask[0, mask_start:non_pad] = True | |
| # If no tokens to predict, mask the last token | |
| if not mask.any(): | |
| mask[0, max(0, non_pad - 1)] = True | |
| # Timestep (low noise -- we want mostly-clean settling) | |
| t = torch.ones(B, dtype=torch.long, device=self.device) | |
| # Embed and forward through lower layers | |
| h_0 = model.embed_input(tokens, t) | |
| h = [h_0] | |
| current = h_0 | |
| msa_start = len(model.forward_blocks) // 2 | |
| max_router_conf = 0.0 | |
| all_retrieved_docs = [] | |
| # Forward pass with MSA retrieval at upper layers | |
| for l, block in enumerate(model.forward_blocks): | |
| current = block(current) | |
| # At MSA layers: retrieve from memory | |
| if l >= msa_start: | |
| mem_k, mem_v, conf, doc_ids = self._retrieve_documents(current, l) | |
| max_router_conf = max(max_router_conf, conf) | |
| all_retrieved_docs.extend(doc_ids) | |
| # Apply sparse attention from MSA layer if we retrieved docs | |
| if mem_k is not None: | |
| msa_idx = l - msa_start | |
| if msa_idx < len(self.msa_layers): | |
| current = self.msa_layers[msa_idx](current, mem_k, mem_v) | |
| h.append(current) | |
| # Store h_init for settling | |
| h_init = [hi.detach() for hi in h] | |
| # --- Settling followed by canonical post-settle learning --- | |
| L = model.n_active_layers | |
| v = [torch.zeros_like(h_init[l + 1]) for l in range(L)] | |
| energies = [] | |
| # Check drift before settling to decide on elastic regularization | |
| current_drift = self.drift_monitor.compute_drift(model, self.msa_layers) | |
| use_elastic = current_drift > config.drift_threshold | |
| # Hard cap: rollback if drift is too large | |
| if current_drift > config.drift_hard_cap and len(self._snapshots) > 1: | |
| self._restore_snapshot(-1) | |
| current_drift = self.drift_monitor.compute_drift(model, self.msa_layers) | |
| for k in range(config.n_settling_steps): | |
| # Adaptive active tokens after first step | |
| if k > 0: | |
| with torch.no_grad(): | |
| uncertainty = model.compute_token_uncertainty(h) | |
| active_tokens = torch.sigmoid( | |
| (uncertainty - mc.settling_threshold) / mc.settling_temperature | |
| ) | |
| else: | |
| active_tokens = None | |
| h, v, energy = model.settling_step( | |
| h, v, h_init, tokens, mask, t, | |
| active_tokens=active_tokens, | |
| ) | |
| energies.append(energy) | |
| model.post_settle_update( | |
| h, | |
| x_input=tokens, | |
| x_0=tokens, | |
| mask=mask, | |
| t=t, | |
| param_lr_scale=config.param_lr_scale, | |
| energies=energies, | |
| ) | |
| # Apply elastic regularization after the canonical model update. | |
| if use_elastic: | |
| penalty = self.drift_monitor.get_elastic_penalty( | |
| model, self.msa_layers, config.elastic_lambda | |
| ) | |
| if penalty.requires_grad: | |
| model.zero_grad(set_to_none=True) | |
| self.msa_layers.zero_grad(set_to_none=True) | |
| penalty.backward() | |
| with torch.no_grad(): | |
| lr = config.param_lr_scale | |
| for p in list(model.parameters()) + list(self.msa_layers.parameters()): | |
| if p.grad is not None: | |
| p.data -= lr * p.grad | |
| p.grad.zero_() | |
| # Also update MSA router parameters from routing errors | |
| self._update_routers(h, tokens, mask, t) | |
| self.energy_traces.append(energies) | |
| # --- Decode answer --- | |
| with torch.no_grad(): | |
| logits = model.readout(model.readout_norm(h[-1])) | |
| probs = F.softmax(logits, dim=-1) | |
| answer_tokens = logits[0, mask_start:non_pad].argmax(dim=-1) | |
| # Compute quality components | |
| energy_ratio = energies[-1] / max(energies[0], 1e-8) if energies else 1.0 | |
| answer_probs = probs[0, mask_start:non_pad] | |
| entropy = -(answer_probs * (answer_probs + 1e-10).log()).sum(dim=-1) | |
| max_entropy = math.log(mc.vocab_size) | |
| answer_coherence = 1.0 - (entropy.mean().item() / max_entropy) | |
| # Record quality | |
| q_n = self.quality_tracker.record( | |
| router_confidence=max_router_conf, | |
| energy_ratio=max(0.0, min(1.0, energy_ratio)), | |
| answer_coherence=answer_coherence, | |
| ) | |
| # Periodic snapshot | |
| if self._query_count % config.snapshot_every == 0: | |
| self._save_snapshot() | |
| return { | |
| "answer_logits": logits, | |
| "answer_tokens": answer_tokens, | |
| "quality": q_n, | |
| "drift": current_drift, | |
| "energy_trace": energies, | |
| "retrieved_docs": list(set(all_retrieved_docs)), | |
| "query_number": self._query_count, | |
| } | |
| def _update_routers(self, h_settled: list, tokens: torch.Tensor, | |
| mask: torch.Tensor, t: torch.Tensor) -> None: | |
| """Update MSA router parameters using settled hidden states. | |
| The router projectors (W_QR, W_KR) are updated via the contrastive | |
| routing loss, using the settled states as signal for what the | |
| "correct" routing should have been. | |
| """ | |
| msa_start = len(self.model.forward_blocks) // 2 | |
| for msa_idx, msa_layer in enumerate(self.msa_layers): | |
| layer_idx = msa_start + msa_idx | |
| if layer_idx + 1 >= len(h_settled): | |
| continue | |
| h_at_layer = h_settled[layer_idx + 1].detach() | |
| # Get routing keys from memory | |
| routing_keys, _ = self.memory_bank.get_routing_keys(layer_idx) | |
| if routing_keys is None: | |
| continue | |
| routing_keys = routing_keys.to(self.device) | |
| # Compute current routing scores | |
| scores = msa_layer.compute_routing_scores(h_at_layer, routing_keys) | |
| # Self-supervised signal: top-scored docs are "positive", | |
| # bottom-scored are "negative" | |
| k = min(self.msa_layers[0].msa_config.top_k, scores.shape[1]) | |
| if scores.shape[1] <= k: | |
| continue | |
| _, top_idx = scores.topk(k, dim=1) | |
| _, bot_idx = scores.topk(scores.shape[1] - k, dim=1, largest=False) | |
| scores_pos = scores.gather(1, top_idx) | |
| scores_neg = scores.gather(1, bot_idx) | |
| # Contrastive loss for router | |
| router_loss = compute_routing_aux_loss( | |
| scores_pos, scores_neg, | |
| temperature=msa_layer.msa_config.aux_temperature, | |
| ) | |
| if router_loss.requires_grad: | |
| router_loss.backward() | |
| lr = self.config.param_lr_scale * self.config.router_lr_boost | |
| with torch.no_grad(): | |
| nn.utils.clip_grad_norm_(msa_layer.router.parameters(), 1.0) | |
| for p in msa_layer.router.parameters(): | |
| if p.grad is not None: | |
| p.data -= lr * p.grad | |
| p.grad.zero_() | |
| # ------------------------------------------------------------------ | |
| # Diagnostics | |
| # ------------------------------------------------------------------ | |
| def get_improvement_curve(self) -> List[float]: | |
| """Return Q_1, Q_2, ..., Q_N quality sequence.""" | |
| return self.quality_tracker.get_improvement_curve() | |
| def get_drift_curve(self) -> List[float]: | |
| """Return drift_1, drift_2, ..., drift_N.""" | |
| return self.drift_monitor.history | |
| def get_diagnostics(self) -> Dict: | |
| """Return comprehensive diagnostics.""" | |
| return { | |
| "n_queries": self._query_count, | |
| "current_quality": self.quality_tracker.current_quality, | |
| "quality_curve": self.get_improvement_curve(), | |
| "quality_rolling": self.quality_tracker.get_rolling_average( | |
| self.config.quality_window | |
| ), | |
| "drift_curve": self.get_drift_curve(), | |
| "drift_components": self.drift_monitor.component_history, | |
| "energy_traces": self.energy_traces, | |
| "n_snapshots": len(self._snapshots), | |
| "memory_bank_size": len(self.memory_bank), | |
| } | |
| def theoretical_bound(self, N: int) -> float: | |
| """Compute the theoretical convergence bound at query N. | |
| E[||W_QR^N - W_QR*||^2] = C / sqrt(N) | |
| The constant C depends on the settling contraction rate rho | |
| and the noise variance sigma^2 of the stochastic gradient: | |
| C = sigma^2 / (1 - rho^K) | |
| where K = n_settling_steps and rho < 1 is the SHO contraction rate. | |
| We estimate C from the empirical quality curve. | |
| """ | |
| if N == 0: | |
| return float("inf") | |
| # Estimate C from the first few queries | |
| if len(self.quality_tracker.history) >= 2: | |
| q1 = 1.0 - self.quality_tracker.history[0] | |
| c_est = q1 # rough: error at N=1 should be ~C/1 | |
| else: | |
| c_est = 1.0 | |
| return c_est / math.sqrt(N) | |
| # ============================================================================ | |
| # Simulation: demonstrate self-improvement over 50 queries | |
| # ============================================================================ | |
| def run_simulation(n_queries: int = 50, device: str = "cpu") -> Dict: | |
| """Run a self-improving retrieval simulation. | |
| Creates a small model, populates a memory bank with synthetic documents, | |
| and issues a sequence of queries. Each query triggers unified settling | |
| that updates both hidden states and retrieval parameters. | |
| Returns: | |
| Dict with quality curve, drift curve, energy traces, and diagnostics. | |
| """ | |
| print("=" * 70) | |
| print("Direction G: Self-Improving Retrieval Simulation") | |
| print("=" * 70) | |
| # --- Setup --- | |
| model_config = PCSHOConfig( | |
| vocab_size=300, | |
| max_seq_len=128, | |
| d_model=128, | |
| n_heads=4, | |
| n_layers=4, | |
| d_ff=256, | |
| n_diffusion_steps=50, | |
| n_settling_steps=4, | |
| eta_base=0.05, | |
| online_learn_lr=1e-4, | |
| ) | |
| msa_config = MSAConfig( | |
| chunk_size=32, | |
| top_k=4, | |
| router_dim=64, | |
| n_router_heads=4, | |
| apply_from_layer=2, | |
| ) | |
| si_config = SelfImprovingConfig( | |
| elastic_lambda=0.005, | |
| drift_threshold=0.15, | |
| drift_hard_cap=0.40, | |
| n_settling_steps=4, | |
| param_lr_scale=0.02, | |
| quality_ema_alpha=0.15, | |
| snapshot_every=10, | |
| ) | |
| print(f"\nModel: d={model_config.d_model}, L={model_config.n_layers}, " | |
| f"H={model_config.n_heads}") | |
| print(f"MSA: top_k={msa_config.top_k}, router_dim={msa_config.router_dim}") | |
| print(f"Settling steps per query: {si_config.n_settling_steps}") | |
| # Create model and MSA layers | |
| model = PCSHODLM(model_config).to(device) | |
| msa_layers = create_msa_layers(model_config, msa_config).to(device) | |
| memory_bank = MemoryBank(chunk_size=msa_config.chunk_size) | |
| param_count = sum(p.numel() for p in model.parameters()) | |
| msa_param_count = sum(p.numel() for p in msa_layers.parameters()) | |
| print(f"Parameters: model={param_count:,}, MSA={msa_param_count:,}") | |
| # --- Populate memory bank with synthetic documents --- | |
| documents = [ | |
| "The speed of light in vacuum is approximately 299792458 meters per second.", | |
| "Photosynthesis converts carbon dioxide and water into glucose and oxygen.", | |
| "The Pythagorean theorem states that a squared plus b squared equals c squared.", | |
| "DNA stores genetic information using four nucleotide bases: A T G and C.", | |
| "Gravity is the force of attraction between objects with mass.", | |
| "Water freezes at zero degrees Celsius and boils at one hundred degrees.", | |
| "The mitochondria are the powerhouse of the cell.", | |
| "Newtons first law states an object in motion stays in motion.", | |
| "The periodic table organizes elements by atomic number and properties.", | |
| "Evolution by natural selection drives adaptation in populations.", | |
| "Quantum mechanics describes behavior of matter at atomic scales.", | |
| "The human genome contains approximately three billion base pairs.", | |
| "Plate tectonics explains the movement of Earths lithospheric plates.", | |
| "Entropy always increases in an isolated system.", | |
| "General relativity describes gravity as curvature of spacetime.", | |
| ] | |
| print(f"\nEncoding {len(documents)} documents into memory bank...") | |
| for i, doc in enumerate(documents): | |
| MemoryEncoder.encode_document( | |
| model, doc, f"doc_{i}", memory_bank, msa_layers, | |
| chunk_size=msa_config.chunk_size, device=device, | |
| ) | |
| print(f"Memory bank: {len(memory_bank)} documents") | |
| # --- Build retriever --- | |
| retriever = SelfImprovingRetriever( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=memory_bank, | |
| config=si_config, | |
| device=device, | |
| ) | |
| # --- Query sequence --- | |
| queries = [ | |
| "What is the speed of light?", | |
| "How do plants make food?", | |
| "What is the Pythagorean theorem?", | |
| "What are the bases of DNA?", | |
| "Why do objects fall?", | |
| "At what temperature does water freeze?", | |
| "What produces energy in cells?", | |
| "What happens to moving objects?", | |
| "How are chemical elements organized?", | |
| "What drives evolution?", | |
| "How do atoms behave?", | |
| "How large is the human genome?", | |
| "What moves the continents?", | |
| "Does entropy increase or decrease?", | |
| "How does gravity work in general relativity?", | |
| # Repeat with variations to show learning | |
| "Tell me about light speed.", | |
| "Explain photosynthesis.", | |
| "Describe the Pythagorean relationship.", | |
| "What nucleotides make up DNA?", | |
| "Why is there gravity?", | |
| "When does water boil?", | |
| "Where is energy made in a cell?", | |
| "Do objects keep moving?", | |
| "What is the periodic table?", | |
| "How does natural selection work?", | |
| "What is quantum mechanics about?", | |
| "How many base pairs in human DNA?", | |
| "What are tectonic plates?", | |
| "Explain the second law of thermodynamics.", | |
| "Describe spacetime curvature.", | |
| # More variations | |
| "Light travels at what speed?", | |
| "CO2 and water become what in plants?", | |
| "Right triangles follow what rule?", | |
| "Adenine thymine guanine cytosine are what?", | |
| "Mass attracts mass through what force?", | |
| "Zero degrees Celsius is the freezing point of what?", | |
| "Mitochondria function is what?", | |
| "Inertia means what?", | |
| "Elements are ordered by what?", | |
| "Survival of the fittest is part of what?", | |
| "Subatomic particles follow what physics?", | |
| "Three billion base pairs are in what?", | |
| "Continental drift is caused by what?", | |
| "Isolated systems and entropy?", | |
| "Einstein described gravity as what?", | |
| # Final batch | |
| "Speed of electromagnetic radiation in vacuum?", | |
| "Chloroplasts perform what process?", | |
| "a^2 + b^2 = c^2 is called what?", | |
| "The double helix stores information using what?", | |
| "What bends spacetime?", | |
| ] | |
| queries = queries[:n_queries] | |
| print(f"\nRunning {len(queries)} queries with self-improving retrieval...\n") | |
| print(f"{'Query':>5} | {'Q_N':>6} | {'Drift':>7} | {'E_ratio':>8} | {'Retrieved':>9} | Text") | |
| print("-" * 90) | |
| for i, q in enumerate(queries): | |
| result = retriever.query(q) | |
| # Energy ratio for display | |
| etrace = result["energy_trace"] | |
| e_ratio = etrace[-1] / max(etrace[0], 1e-8) if len(etrace) >= 2 else 1.0 | |
| print(f"{result['query_number']:>5} | {result['quality']:>6.3f} | " | |
| f"{result['drift']:>7.4f} | {e_ratio:>8.4f} | " | |
| f"{len(result['retrieved_docs']):>9} | {q[:40]}") | |
| # --- Summary --- | |
| diagnostics = retriever.get_diagnostics() | |
| curve = diagnostics["quality_curve"] | |
| drift = diagnostics["drift_curve"] | |
| print("\n" + "=" * 70) | |
| print("RESULTS SUMMARY") | |
| print("=" * 70) | |
| # Quality improvement | |
| first_5 = sum(curve[:5]) / min(5, len(curve)) | |
| last_5 = sum(curve[-5:]) / min(5, len(curve)) | |
| print(f"\nRetrieval Quality (Q_N):") | |
| print(f" First 5 queries (avg): {first_5:.4f}") | |
| print(f" Last 5 queries (avg): {last_5:.4f}") | |
| print(f" Improvement: {last_5 - first_5:+.4f} ({(last_5/max(first_5,1e-8) - 1)*100:+.1f}%)") | |
| print(f" Final EMA quality: {diagnostics['current_quality']:.4f}") | |
| # Drift | |
| if drift: | |
| print(f"\nParameter Drift:") | |
| print(f" Final drift: {drift[-1]:.4f}") | |
| print(f" Max drift: {max(drift):.4f}") | |
| print(f" Elastic reg activated: {sum(1 for d in drift if d > si_config.drift_threshold)} times") | |
| # Convergence bound | |
| print(f"\nConvergence Bound E[||W_QR^N - W_QR*||^2] = O(1/sqrt(N)):") | |
| for n in [1, 10, 25, 50]: | |
| if n <= n_queries: | |
| bound = retriever.theoretical_bound(n) | |
| actual = 1.0 - (curve[n - 1] if n <= len(curve) else curve[-1]) | |
| print(f" N={n:>3}: bound={bound:.4f}, actual_error={actual:.4f}") | |
| print(f"\nSnapshots saved: {diagnostics['n_snapshots']}") | |
| print(f"Memory bank: {diagnostics['memory_bank_size']} documents") | |
| # ASCII quality curve | |
| print(f"\nQuality Curve (Q_N over queries):") | |
| rolling = diagnostics["quality_rolling"] | |
| if rolling: | |
| max_q = max(rolling) if max(rolling) > 0 else 1.0 | |
| min_q = min(rolling) | |
| bar_width = 40 | |
| for i, q in enumerate(rolling): | |
| if i % max(1, len(rolling) // 20) == 0 or i == len(rolling) - 1: | |
| normalized = (q - min_q) / max(max_q - min_q, 1e-8) | |
| bar = "#" * int(normalized * bar_width) | |
| print(f" Q_{i+1:>3}: {q:.3f} |{bar}") | |
| return diagnostics | |
| # ============================================================================ | |
| # Entry point | |
| # ============================================================================ | |
| if __name__ == "__main__": | |
| diagnostics = run_simulation(n_queries=50, device="cpu") | |