pc-sho-dlm-code / src /multihop.py
zotowata's picture
Sync 2B training job files
2a5274f verified
Raw History Blame Contribute Delete
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
# ──────────────────────────────────────────────────────────────────────
@dataclass
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
@dataclass
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()