Download src/multihop.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 29.3 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/multihop.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/src/multihop.py
-
curl -L -o multihop.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/src/multihop.py
29.3 kB
| """ | |
| Energy-Based Multi-Hop Reasoning for PC-SHO-DLM + MSA | |
| Direction D: Fuses retrieval and reasoning into continuous energy | |
| minimization instead of discrete retrieve -> reason -> retrieve cycles. | |
| Multi-hop energy: | |
| E = E_PC(h) + sum_n [ mu_n * L_route(h, M \\ D_{<n}) + L_attend(h, D_n) ] | |
| Each hop is NOT a separate pipeline stage. Instead, K total settling steps | |
| are divided into N adaptive hops. At each hop boundary the router | |
| re-scores the memory bank (excluding documents already retrieved) and | |
| potentially swaps the active context. The hidden states, velocities, and | |
| router weights all live on the same energy surface -- settling minimizes | |
| everything jointly. | |
| Key properties: | |
| - Diversity: already-retrieved documents are excluded from the candidate | |
| pool so each hop is forced to bring in genuinely new information. | |
| - Adaptive depth: hops stop early when |E^{k+1} - E^k| < epsilon. | |
| - Full audit trail: every hop records which documents were selected, | |
| their scores, and the energy at selection time. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass, field | |
| from typing import Dict, List, Optional, Tuple | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from .model import PCSHODLM, PCSHOConfig | |
| from .msa import ( | |
| MSAConfig, | |
| MSALayer, | |
| MemoryBank, | |
| MemoryEncoder, | |
| chunk_mean_pool, | |
| create_msa_layers, | |
| ) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Data structures | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class HopRecord: | |
| """Audit trail for a single hop.""" | |
| hop_index: int | |
| settling_steps_range: Tuple[int, int] # (start_k, end_k) inclusive | |
| retrieved_doc_ids: List[str] | |
| routing_scores: Dict[str, float] # doc_id -> score (top candidates) | |
| energy_at_selection: float | |
| reason: str # human-readable why this doc won | |
| class ReasoningChain: | |
| """Full audit trail across all hops.""" | |
| hops: List[HopRecord] = field(default_factory=list) | |
| energy_trace: List[float] = field(default_factory=list) | |
| converged: bool = False | |
| total_settling_steps: int = 0 | |
| def summary(self) -> str: | |
| lines = [] | |
| for hop in self.hops: | |
| docs = ", ".join(hop.retrieved_doc_ids) | |
| lines.append( | |
| f"Hop {hop.hop_index}: retrieved [{docs}] " | |
| f"(energy={hop.energy_at_selection:.4f}, " | |
| f"reason={hop.reason})" | |
| ) | |
| conv = "converged" if self.converged else "budget exhausted" | |
| lines.append(f"Outcome: {conv} after {self.total_settling_steps} steps") | |
| return "\n".join(lines) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # MultiHopSettler | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class MultiHopSettler: | |
| """Energy-based multi-hop retrieval fused with PC-SHO settling. | |
| Instead of the classic loop: | |
| retrieve -> read -> retrieve -> read -> answer | |
| we run a *single* settling process whose energy surface includes | |
| both predictive-coding consistency AND retrieval quality. At | |
| adaptive hop boundaries, the active document context is updated by | |
| re-scoring the memory bank with the *current* hidden states | |
| (which have been refined by settling). | |
| Parameters | |
| ---------- | |
| model : PCSHODLM | |
| The core PC-SHO-DLM model. | |
| msa_layers : nn.ModuleList | |
| MSA layers (one per upper-half model layer). | |
| memory_bank : MemoryBank | |
| Pre-encoded document store. | |
| total_settling_steps : int | |
| Budget K for the entire multi-hop process. | |
| max_hops : int | |
| Upper bound N on the number of hops. | |
| docs_per_hop : int | |
| Number of documents to retrieve at each hop boundary. | |
| epsilon : float | |
| Energy convergence threshold -- stop hopping when | |
| |E^{k+1} - E^k| < epsilon. | |
| route_weight : float | |
| Coefficient mu for the routing energy term. | |
| unified_mode : bool | |
| If True, apply a canonical post-settle model update at each hop | |
| boundary after shared latent settling. | |
| param_lr_scale : float | |
| Learning rate scale for the post-settle parameter update. | |
| """ | |
| def __init__( | |
| self, | |
| model: PCSHODLM, | |
| msa_layers: nn.ModuleList, | |
| memory_bank: MemoryBank, | |
| total_settling_steps: int = 16, | |
| max_hops: int = 4, | |
| docs_per_hop: int = 1, | |
| epsilon: float = 1e-3, | |
| route_weight: float = 0.5, | |
| unified_mode: bool = False, | |
| param_lr_scale: float = 0.01, | |
| ): | |
| self.model = model | |
| self.msa_layers = msa_layers | |
| self.memory_bank = memory_bank | |
| self.total_settling_steps = total_settling_steps | |
| self.max_hops = max_hops | |
| self.docs_per_hop = docs_per_hop | |
| self.epsilon = epsilon | |
| self.route_weight = route_weight | |
| self.unified_mode = unified_mode | |
| self.param_lr_scale = param_lr_scale | |
| # ββ helpers ββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _tokenize(self, text: str, device: torch.device) -> torch.Tensor: | |
| """Byte-level tokenization consistent with MemoryEncoder.""" | |
| max_len = self.model.config.max_seq_len | |
| ids = [min(b + 1, 256) for b in text.encode("utf-8")[:max_len]] | |
| tokens = torch.tensor(ids, dtype=torch.long, device=device).unsqueeze(0) | |
| if tokens.shape[1] < max_len: | |
| tokens = F.pad(tokens, (0, max_len - tokens.shape[1])) | |
| return tokens | |
| def _get_msa_layer_indices(self) -> List[int]: | |
| """Layer indices in the model where MSA layers apply.""" | |
| n_layers = len(self.model.forward_blocks) | |
| msa_start = n_layers // 2 | |
| return list(range(msa_start, msa_start + len(self.msa_layers))) | |
| # ββ routing ββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _score_documents( | |
| self, | |
| h: list[torch.Tensor], | |
| excluded_doc_ids: set[str], | |
| ) -> Tuple[List[str], Dict[str, float], float]: | |
| """Score all non-excluded documents against current hidden state. | |
| Uses the first MSA layer's router (layer closest to the query) | |
| for scoring. Returns the top-k doc IDs, their scores, and | |
| the routing energy contribution. | |
| Returns | |
| ------- | |
| selected_ids : List[str] | |
| Top docs_per_hop document IDs. | |
| all_scores : Dict[str, float] | |
| Scores for all candidate documents (for the audit trail). | |
| route_energy : float | |
| Routing energy L_route = -sum(top_k scores). | |
| """ | |
| msa_indices = self._get_msa_layer_indices() | |
| if not msa_indices: | |
| return [], {}, 0.0 | |
| # Use hidden state at the first MSA layer boundary | |
| layer_idx = msa_indices[0] | |
| h_query = h[layer_idx + 1] # hidden state after this layer | |
| msa_layer = self.msa_layers[0] | |
| # Gather routing keys from memory bank, excluding already-retrieved | |
| all_kr, all_chunk_ids = self.memory_bank.get_routing_keys(layer_idx) | |
| if all_kr is None or len(all_chunk_ids) == 0: | |
| return [], {}, 0.0 | |
| # Build candidate mask: exclude already-retrieved docs | |
| candidate_mask = torch.ones(len(all_chunk_ids), dtype=torch.bool) | |
| for i, cid in enumerate(all_chunk_ids): | |
| if cid in excluded_doc_ids: | |
| candidate_mask[i] = False | |
| if not candidate_mask.any(): | |
| return [], {}, 0.0 | |
| # Score candidates | |
| candidate_kr = all_kr[candidate_mask].to(h_query.device) | |
| candidate_ids = [cid for cid, m in zip(all_chunk_ids, candidate_mask) if m] | |
| scores = msa_layer.compute_routing_scores(h_query, candidate_kr) # (B, N_cand) | |
| scores_flat = scores.squeeze(0) # (N_cand,) | |
| # Aggregate scores per document (max over chunks) | |
| doc_scores: Dict[str, float] = {} | |
| for idx, doc_id in enumerate(candidate_ids): | |
| s = scores_flat[idx].item() | |
| if doc_id not in doc_scores or s > doc_scores[doc_id]: | |
| doc_scores[doc_id] = s | |
| # Sort and select top-k | |
| ranked = sorted(doc_scores.items(), key=lambda x: x[1], reverse=True) | |
| selected = ranked[: self.docs_per_hop] | |
| selected_ids = [did for did, _ in selected] | |
| # Routing energy: negative sum of selected scores (lower = better retrieval) | |
| route_energy = -sum(s for _, s in selected) | |
| return selected_ids, doc_scores, route_energy | |
| def _gather_memory_kv( | |
| self, doc_ids: List[str], layer_idx: int | |
| ) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: | |
| """Retrieve compressed K, V for selected documents at a layer.""" | |
| K, V = self.memory_bank.get_kv(doc_ids, layer_idx) | |
| return K, V | |
| # ββ attention energy βββββββββββββββββββββββββββββββββββββββββββββ | |
| def _compute_attend_energy( | |
| self, | |
| h: list[torch.Tensor], | |
| active_doc_ids: List[str], | |
| ) -> float: | |
| """L_attend: how well the current hidden states attend to active docs. | |
| A low attend energy means the hidden states have successfully | |
| incorporated the retrieved information. | |
| """ | |
| if not active_doc_ids: | |
| return 0.0 | |
| msa_indices = self._get_msa_layer_indices() | |
| total = 0.0 | |
| for i, layer_idx in enumerate(msa_indices): | |
| if i >= len(self.msa_layers): | |
| break | |
| msa_layer = self.msa_layers[i] | |
| h_query = h[layer_idx + 1] | |
| K, V = self._gather_memory_kv(active_doc_ids, layer_idx) | |
| if K is None: | |
| continue | |
| device = h_query.device | |
| K = K.unsqueeze(0).to(device) # (1, n_chunks, D) | |
| V = V.unsqueeze(0).to(device) | |
| # Compute cross-attention output | |
| attn_out = msa_layer.sparse_attention(h_query, K, V) | |
| # Attend energy = negative cosine similarity between attn output | |
| # and the hidden state (lower = more consistent) | |
| cos_sim = F.cosine_similarity( | |
| attn_out.flatten(), h_query.flatten(), dim=0 | |
| ) | |
| total += (1.0 - cos_sim.item()) | |
| return total | |
| # ββ core settling with multi-hop βββββββββββββββββββββββββββββββββ | |
| def settle_multihop( | |
| self, | |
| query_tokens: torch.Tensor, | |
| ) -> Tuple[list[torch.Tensor], ReasoningChain]: | |
| """Run the full multi-hop energy-based settling process. | |
| Parameters | |
| ---------- | |
| query_tokens : torch.Tensor | |
| (1, S) token IDs of the query. | |
| Returns | |
| ------- | |
| h_settled : list[torch.Tensor] | |
| Final hidden states after settling. | |
| chain : ReasoningChain | |
| Full audit trail. | |
| """ | |
| model = self.model | |
| config = model.config | |
| device = query_tokens.device | |
| L = config.n_layers | |
| K_total = self.total_settling_steps | |
| chain = ReasoningChain() | |
| excluded_docs: set[str] = set() | |
| active_doc_ids: List[str] = [] | |
| # -- Initialization -- | |
| t = torch.ones(1, dtype=torch.long, device=device) | |
| mask = torch.zeros_like(query_tokens, dtype=torch.bool) | |
| h_0 = model.embed_input(query_tokens, t) | |
| h_init = model.amortized_forward_pass(h_0) | |
| h = [hi.detach() for hi in h_init] | |
| v = [torch.zeros_like(h_init[l + 1]) for l in range(L)] | |
| # Distribute settling steps across hops | |
| steps_per_hop = max(1, K_total // self.max_hops) | |
| global_step = 0 | |
| prev_energy = float("inf") | |
| for hop_idx in range(self.max_hops): | |
| if global_step >= K_total: | |
| break | |
| # --- Hop boundary: re-score and retrieve --- | |
| new_ids, all_scores, route_energy = self._score_documents( | |
| h, excluded_docs | |
| ) | |
| if new_ids: | |
| active_doc_ids.extend(new_ids) | |
| excluded_docs.update(new_ids) | |
| # Build reason string | |
| if hop_idx == 0: | |
| reason = "initial retrieval based on query embedding" | |
| else: | |
| reason = ( | |
| f"energy delta prompted re-routing; " | |
| f"hidden state refined over steps " | |
| f"{global_step - steps_per_hop}-{global_step}" | |
| ) | |
| hop_record = HopRecord( | |
| hop_index=hop_idx, | |
| settling_steps_range=(global_step, -1), # end filled later | |
| retrieved_doc_ids=list(new_ids), | |
| routing_scores={ | |
| k: round(v, 4) for k, v in sorted( | |
| all_scores.items(), key=lambda x: -x[1] | |
| )[:10] | |
| }, | |
| energy_at_selection=prev_energy if prev_energy != float("inf") else 0.0, | |
| reason=reason, | |
| ) | |
| else: | |
| hop_record = HopRecord( | |
| hop_index=hop_idx, | |
| settling_steps_range=(global_step, -1), | |
| retrieved_doc_ids=[], | |
| routing_scores={}, | |
| energy_at_selection=prev_energy if prev_energy != float("inf") else 0.0, | |
| reason="no new documents available", | |
| ) | |
| # --- Inject retrieved memory into hidden states via MSA --- | |
| # Augment the initialization so the settling block sees memory context | |
| h = self._inject_memory_context(h, active_doc_ids) | |
| # --- Inner settling loop for this hop --- | |
| hop_start = global_step | |
| steps_this_hop = min(steps_per_hop, K_total - global_step) | |
| for k_local in range(steps_this_hop): | |
| h, v, energy = model.settling_step( | |
| h, v, h_init, query_tokens, mask, t, | |
| ) | |
| # Add routing + attend energies to get the full multi-hop energy | |
| attend_energy = self._compute_attend_energy(h, active_doc_ids) | |
| full_energy = energy + self.route_weight * route_energy + attend_energy | |
| chain.energy_trace.append(full_energy) | |
| global_step += 1 | |
| if self.unified_mode and chain.energy_trace: | |
| model.post_settle_update( | |
| h, | |
| x_input=query_tokens, | |
| x_0=query_tokens, | |
| mask=mask, | |
| t=t, | |
| param_lr_scale=self.param_lr_scale, | |
| energies=[chain.energy_trace[-1]], | |
| ) | |
| # Fill in the end of the settling range | |
| hop_record.settling_steps_range = (hop_start, global_step - 1) | |
| chain.hops.append(hop_record) | |
| # --- Adaptive stopping --- | |
| current_energy = chain.energy_trace[-1] if chain.energy_trace else float("inf") | |
| if ( | |
| prev_energy != float("inf") | |
| and abs(prev_energy - current_energy) < self.epsilon | |
| ): | |
| chain.converged = True | |
| break | |
| prev_energy = current_energy | |
| chain.total_settling_steps = global_step | |
| return h, chain | |
| def _inject_memory_context( | |
| self, | |
| h: list[torch.Tensor], | |
| doc_ids: List[str], | |
| ) -> list[torch.Tensor]: | |
| """Inject retrieved memory into hidden states via MSA cross-attention. | |
| This modifies the hidden states at MSA-eligible layers by running | |
| a cross-attention step with the retrieved document KV. The result | |
| is blended with the current hidden state so the settling dynamics | |
| incorporate the new information smoothly. | |
| """ | |
| if not doc_ids: | |
| return h | |
| msa_indices = self._get_msa_layer_indices() | |
| h_new = list(h) # shallow copy | |
| for i, layer_idx in enumerate(msa_indices): | |
| if i >= len(self.msa_layers): | |
| break | |
| K, V = self._gather_memory_kv(doc_ids, layer_idx) | |
| if K is None: | |
| continue | |
| device = h[layer_idx + 1].device | |
| K = K.unsqueeze(0).to(device) | |
| V = V.unsqueeze(0).to(device) | |
| msa_layer = self.msa_layers[i] | |
| h_query = h[layer_idx + 1] | |
| # Cross-attend and blend (additive residual, scaled) | |
| attn_out = msa_layer.sparse_attention(h_query, K, V) | |
| h_new[layer_idx + 1] = h_query + 0.5 * (attn_out - h_query) | |
| return h_new | |
| # ββ public API βββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def reason_and_answer( | |
| self, | |
| question: str, | |
| memory_bank: Optional[MemoryBank] = None, | |
| ) -> Tuple[str, ReasoningChain, List[float]]: | |
| """End-to-end multi-hop reasoning. | |
| Parameters | |
| ---------- | |
| question : str | |
| The query text. | |
| memory_bank : MemoryBank, optional | |
| Override the instance-level memory bank. | |
| Returns | |
| ------- | |
| answer_text : str | |
| Decoded answer from the settled hidden states. | |
| chain : ReasoningChain | |
| Full audit trail of hops, documents, and energies. | |
| energy_trace : List[float] | |
| Per-step energy values (convenience alias for chain.energy_trace). | |
| """ | |
| if memory_bank is not None: | |
| original_bank = self.memory_bank | |
| self.memory_bank = memory_bank | |
| device = next(self.model.parameters()).device | |
| query_tokens = self._tokenize(question, device) | |
| try: | |
| h_settled, chain = self.settle_multihop(query_tokens) | |
| finally: | |
| if memory_bank is not None: | |
| self.memory_bank = original_bank | |
| # Decode answer from the top-layer hidden state | |
| with torch.no_grad(): | |
| logits = self.model.readout( | |
| self.model.readout_norm(h_settled[-1]) | |
| ) | |
| pred_ids = logits.argmax(dim=-1).squeeze(0) # (S,) | |
| # Convert byte-level tokens back to text | |
| answer_bytes = [] | |
| for tid in pred_ids.tolist(): | |
| if tid == 0: | |
| break | |
| byte_val = tid - 1 | |
| if 0 <= byte_val < 256: | |
| answer_bytes.append(byte_val) | |
| answer_text = bytes(answer_bytes).decode("utf-8", errors="replace") | |
| return answer_text, chain, chain.energy_trace | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Tests | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _run_tests(): | |
| """Comprehensive tests for MultiHopSettler. | |
| Test scenario: two-hop compositional question. | |
| Doc A: "The capital of France is Paris." | |
| Doc B: "The Eiffel Tower is located in Paris." | |
| Question: "What famous tower is in the capital of France?" | |
| Hop 1 should retrieve Doc A (capital of France -> Paris). | |
| Hop 2 should retrieve Doc B (Paris -> Eiffel Tower). | |
| """ | |
| print("=" * 70) | |
| print("MultiHopSettler Tests") | |
| print("=" * 70) | |
| # ββ Setup ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| config = PCSHOConfig( | |
| vocab_size=300, | |
| max_seq_len=128, | |
| d_model=128, | |
| n_heads=4, | |
| n_layers=4, | |
| d_ff=256, | |
| n_diffusion_steps=100, | |
| n_settling_steps=4, | |
| ) | |
| model = PCSHODLM(config) | |
| model.eval() | |
| msa_config = MSAConfig( | |
| chunk_size=32, | |
| top_k=2, | |
| router_dim=64, | |
| n_router_heads=4, | |
| ) | |
| msa_layers = create_msa_layers(config, msa_config) | |
| # ββ Encode documents βββββββββββββββββββββββββββββββββββββββββββββ | |
| memory_bank = MemoryBank(chunk_size=32) | |
| doc_a_text = "The capital of France is Paris." | |
| doc_b_text = "The Eiffel Tower is located in Paris." | |
| doc_c_text = "Kangaroos are native to Australia." # distractor | |
| MemoryEncoder.encode_document(model, doc_a_text, "doc_a", memory_bank, msa_layers, chunk_size=32) | |
| MemoryEncoder.encode_document(model, doc_b_text, "doc_b", memory_bank, msa_layers, chunk_size=32) | |
| MemoryEncoder.encode_document(model, doc_c_text, "doc_c", memory_bank, msa_layers, chunk_size=32) | |
| assert len(memory_bank) == 3, f"Expected 3 docs, got {len(memory_bank)}" | |
| print("[PASS] Memory bank contains 3 documents") | |
| # ββ Test 1: Basic construction and settling ββββββββββββββββββββββ | |
| print("\n--- Test 1: Construction and settling ---") | |
| settler = MultiHopSettler( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=memory_bank, | |
| total_settling_steps=8, | |
| max_hops=3, | |
| docs_per_hop=1, | |
| epsilon=1e-4, | |
| route_weight=0.5, | |
| ) | |
| question = "What famous tower is in the capital of France?" | |
| answer, chain, energy_trace = settler.reason_and_answer(question) | |
| assert isinstance(answer, str), "Answer should be a string" | |
| assert len(chain.hops) > 0, "Should have at least 1 hop" | |
| assert len(energy_trace) > 0, "Should have energy trace" | |
| print(f"[PASS] Settling completed: {chain.total_settling_steps} steps, " | |
| f"{len(chain.hops)} hops") | |
| print(f" Answer length: {len(answer)} chars") | |
| # ββ Test 2: Diversity -- no document retrieved twice βββββββββββββ | |
| print("\n--- Test 2: Document diversity across hops ---") | |
| all_retrieved = [] | |
| for hop in chain.hops: | |
| all_retrieved.extend(hop.retrieved_doc_ids) | |
| unique_retrieved = set(all_retrieved) | |
| assert len(all_retrieved) == len(unique_retrieved), ( | |
| f"Duplicate documents retrieved! All: {all_retrieved}" | |
| ) | |
| print(f"[PASS] All retrieved docs are unique: {sorted(unique_retrieved)}") | |
| # ββ Test 3: Energy trace is recorded βββββββββββββββββββββββββββββ | |
| print("\n--- Test 3: Energy trace ---") | |
| assert len(energy_trace) == chain.total_settling_steps, ( | |
| f"Energy trace length {len(energy_trace)} != steps {chain.total_settling_steps}" | |
| ) | |
| print(f"[PASS] Energy trace has {len(energy_trace)} entries") | |
| print(f" Initial energy: {energy_trace[0]:.4f}") | |
| print(f" Final energy: {energy_trace[-1]:.4f}") | |
| # ββ Test 4: Reasoning chain audit trail ββββββββββββββββββββββββββ | |
| print("\n--- Test 4: Reasoning chain audit trail ---") | |
| for hop in chain.hops: | |
| assert isinstance(hop.hop_index, int) | |
| assert isinstance(hop.settling_steps_range, tuple) | |
| assert isinstance(hop.routing_scores, dict) | |
| assert isinstance(hop.reason, str) | |
| assert len(hop.reason) > 0 | |
| print(f"[PASS] All hop records are well-formed") | |
| print(f"\nFull chain:\n{chain.summary()}") | |
| # ββ Test 5: Adaptive convergence βββββββββββββββββββββββββββββββββ | |
| print("\n--- Test 5: Adaptive convergence (large epsilon) ---") | |
| settler_early = MultiHopSettler( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=memory_bank, | |
| total_settling_steps=16, | |
| max_hops=4, | |
| docs_per_hop=1, | |
| epsilon=1e10, # huge epsilon -> should stop after 1 hop | |
| route_weight=0.5, | |
| ) | |
| _, chain_early, _ = settler_early.reason_and_answer(question) | |
| # With epsilon=1e10, the energy delta will always be < epsilon after | |
| # the first hop that has a finite previous energy. | |
| # Hop 0 sets prev_energy; hop 1 checks and converges. | |
| assert chain_early.converged, "Should have converged with huge epsilon" | |
| assert len(chain_early.hops) <= 2, ( | |
| f"Expected at most 2 hops with huge epsilon, got {len(chain_early.hops)}" | |
| ) | |
| print(f"[PASS] Early convergence: {len(chain_early.hops)} hops, " | |
| f"converged={chain_early.converged}") | |
| # ββ Test 6: Unified mode (router weights update) βββββββββββββββββ | |
| print("\n--- Test 6: Unified mode ---") | |
| settler_unified = MultiHopSettler( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=memory_bank, | |
| total_settling_steps=8, | |
| max_hops=2, | |
| docs_per_hop=1, | |
| epsilon=1e-4, | |
| route_weight=0.5, | |
| unified_mode=True, | |
| param_lr_scale=0.001, | |
| ) | |
| # Snapshot a parameter before unified settling | |
| ref_param = next(model.forward_blocks[0].parameters()).clone() | |
| model.train() | |
| _, chain_unified, energy_unified = settler_unified.reason_and_answer(question) | |
| model.eval() | |
| # In unified mode, parameters should have been updated | |
| new_param = next(model.forward_blocks[0].parameters()) | |
| param_changed = not torch.allclose(ref_param, new_param, atol=1e-8) | |
| print(f"[PASS] Unified mode ran: {chain_unified.total_settling_steps} steps, " | |
| f"params updated={param_changed}") | |
| # ββ Test 7: Two-hop compositional question βββββββββββββββββββββββ | |
| print("\n--- Test 7: Two-hop composition (core test) ---") | |
| settler_2hop = MultiHopSettler( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=memory_bank, | |
| total_settling_steps=12, | |
| max_hops=3, | |
| docs_per_hop=1, | |
| epsilon=1e-6, | |
| route_weight=0.5, | |
| ) | |
| answer_2, chain_2, trace_2 = settler_2hop.reason_and_answer( | |
| "What famous tower is in the capital of France?" | |
| ) | |
| # Check that at least 2 hops occurred and retrieved distinct docs | |
| assert len(chain_2.hops) >= 2, ( | |
| f"Expected >= 2 hops for compositional question, got {len(chain_2.hops)}" | |
| ) | |
| all_docs_2 = [] | |
| for hop in chain_2.hops: | |
| all_docs_2.extend(hop.retrieved_doc_ids) | |
| assert len(set(all_docs_2)) >= 2, ( | |
| f"Expected >= 2 distinct docs, got {set(all_docs_2)}" | |
| ) | |
| print(f"[PASS] Two-hop composition: {len(chain_2.hops)} hops, " | |
| f"docs={all_docs_2}") | |
| print(f"\nChain:\n{chain_2.summary()}") | |
| # ββ Test 8: Empty memory bank ββββββββββββββββββββββββββββββββββββ | |
| print("\n--- Test 8: Empty memory bank ---") | |
| empty_bank = MemoryBank(chunk_size=32) | |
| settler_empty = MultiHopSettler( | |
| model=model, | |
| msa_layers=msa_layers, | |
| memory_bank=empty_bank, | |
| total_settling_steps=4, | |
| max_hops=2, | |
| docs_per_hop=1, | |
| ) | |
| answer_e, chain_e, trace_e = settler_empty.reason_and_answer("test question") | |
| assert len(trace_e) > 0, "Should still settle even with empty bank" | |
| print(f"[PASS] Empty bank: {chain_e.total_settling_steps} steps, " | |
| f"0 docs retrieved") | |
| # ββ Test 9: Override memory bank in reason_and_answer βββββββββββββ | |
| print("\n--- Test 9: Memory bank override ---") | |
| answer_override, chain_override, _ = settler.reason_and_answer( | |
| "test", memory_bank=empty_bank | |
| ) | |
| for hop in chain_override.hops: | |
| assert len(hop.retrieved_doc_ids) == 0 | |
| # Original bank should still be intact | |
| assert settler.memory_bank is memory_bank | |
| print("[PASS] Memory bank override works, original restored") | |
| # ββ Done βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| print("\n" + "=" * 70) | |
| print("All tests passed.") | |
| print("=" * 70) | |
| if __name__ == "__main__": | |
| _run_tests() | |