File size: 37,107 Bytes
49589ef
 
 
 
2a5274f
 
 
 
49589ef
 
 
 
2a5274f
 
 
 
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a5274f
 
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a5274f
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a5274f
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a5274f
49589ef
 
 
 
 
2a5274f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a5274f
49589ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
"""
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
# ============================================================================

@dataclass
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

    @property
    def current_quality(self) -> float:
        return self._ema if self._initialized else 0.0

    @property
    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")