""" PC-SHO-DLM Online Learning & Dynamic Architecture Evaluation Tests the unique PC capabilities: 1. Online learning during inference (domain adaptation at generation time) 2. Dynamic layer insertion with online learning convergence 3. Comparison: generation quality with/without online learning """ import csv import json import os import sys import time from pathlib import Path import torch import torch.nn.functional as F sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) from model import ( PCSHODLM, PCSHOConfig, InferenceUpdater, DynamicLayerManager, count_parameters ) def load_ao3_text(data_dir: str, max_rows: int = 5000) -> str: """Extract text content from AO3 metadata CSV for domain adaptation testing.""" csv_path = os.path.join(data_dir, "random works - Oct 2025.csv") if not os.path.exists(csv_path): # Try alternative for f in os.listdir(data_dir): if f.endswith(".csv") and "random works" in f: csv_path = os.path.join(data_dir, f) break texts = [] with open(csv_path, "r", errors="replace") as f: reader = csv.DictReader(f) for i, row in enumerate(reader): if i >= max_rows: break # Build a text representation from metadata parts = [] if row.get("Title"): parts.append(f"Title: {row['Title']}") if row.get("Fandom Tags"): parts.append(f"Fandom: {row['Fandom Tags']}") if row.get("Summary Paragraphs"): summary = row["Summary Paragraphs"] if summary and summary != "0": parts.append(f"Summary: {summary}") if row.get("Freeform Tags"): parts.append(f"Tags: {row['Freeform Tags']}") if parts: texts.append(" | ".join(parts)) return "\n".join(texts) def encode_text(text: str, seq_len: int, vocab_size: int = 257) -> torch.Tensor: """Encode text to byte-level tensor (matching training encoding).""" data = torch.tensor( [min(b + 1, vocab_size - 1) for b in text.encode("utf-8")[:seq_len]], dtype=torch.long, ) # Pad if needed if len(data) < seq_len: data = F.pad(data, (0, seq_len - len(data)), value=0) return data.unsqueeze(0) # (1, seq_len) # ============================================================================= # Experiment 1: Online Learning Domain Adaptation # ============================================================================= def test_online_learning_adaptation( checkpoint_path: str, ao3_data_dir: str, device: str = "cpu", ) -> dict: """Test online learning: train on WikiText, adapt to AO3 style at inference. Compares generation energy with and without online learning when the model encounters out-of-domain text. """ print("=" * 60) print("Experiment 1: Online Learning Domain Adaptation") print("=" * 60) # Load trained model ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False) config = ckpt["config"] model = PCSHODLM(config).to(device) model.load_state_dict(ckpt["model_state_dict"]) model.eval() print(f"Loaded model from {checkpoint_path}") # Load AO3 text as "context" for adaptation ao3_text = load_ao3_text(ao3_data_dir, max_rows=1000) print(f"Loaded {len(ao3_text):,} chars of AO3 text") seq_len = config.max_seq_len # Measure baseline energy on AO3 text (no online learning) print("\nBaseline (no online learning):") x_ao3 = encode_text(ao3_text, seq_len, config.vocab_size).to(device) t = torch.full((1,), config.n_diffusion_steps // 2, device=device, dtype=torch.long) x_t, mask = model.schedule.corrupt(x_ao3, t, config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) with torch.no_grad(): h_settled, _, energies_baseline, _, _ = model.settle(h_init, x_ao3, mask, t) logits = model.readout(model.readout_norm(h_settled[-1])) mask_logits = logits[mask] mask_targets = x_ao3[mask] if mask_logits.numel() > 0: loss_baseline = F.cross_entropy(mask_logits, mask_targets).item() else: loss_baseline = 0.0 print(f" Loss: {loss_baseline:.4f}") print(f" Energy: {energies_baseline[-1]:.0f}") # Now generate with online learning enabled print("\nWith online learning (adapting during generation):") # Save initial params param_snapshot = {n: p.data.clone() for n, p in model.named_parameters()} gen_online = model.generate( seq_len=seq_len, batch_size=1, device=device, online_learn=True ) # Measure energy on AO3 text AFTER online adaptation x_t2, mask2 = model.schedule.corrupt(x_ao3, t, config.mask_token_id) h_02 = model.embed_input(x_t2, t) h_init2 = model.amortized_forward_pass(h_02) with torch.no_grad(): h_settled2, _, energies_adapted, _, _ = model.settle(h_init2, x_ao3, mask2, t) logits2 = model.readout(model.readout_norm(h_settled2[-1])) mask_logits2 = logits2[mask2] mask_targets2 = x_ao3[mask2] if mask_logits2.numel() > 0: loss_adapted = F.cross_entropy(mask_logits2, mask_targets2).item() else: loss_adapted = 0.0 # Measure parameter drift total_drift = 0.0 n_params = 0 for n, p in model.named_parameters(): if n in param_snapshot: total_drift += (p.data - param_snapshot[n]).norm().item() n_params += p.numel() print(f" Loss after adaptation: {loss_adapted:.4f}") print(f" Energy after adaptation: {energies_adapted[-1]:.0f}") print(f" Parameter drift: {total_drift:.6f}") print(f" Loss improvement: {loss_baseline - loss_adapted:.4f}") # Restore original params for n, p in model.named_parameters(): if n in param_snapshot: p.data.copy_(param_snapshot[n]) results = { "loss_baseline": loss_baseline, "loss_adapted": loss_adapted, "energy_baseline": energies_baseline[-1], "energy_adapted": energies_adapted[-1], "param_drift": total_drift, "improvement": loss_baseline - loss_adapted, } return results # ============================================================================= # Experiment 2: Dynamic Layer Insertion # ============================================================================= def test_dynamic_layer_insertion( checkpoint_path: str, device: str = "cpu", ) -> dict: """Test inserting layers at runtime and measuring convergence. Inserts 2 new layers into a trained model and measures: - How many settling steps before energy improves - Loss before and after insertion + settling """ print("\n" + "=" * 60) print("Experiment 2: Dynamic Layer Insertion") print("=" * 60) ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False) config = ckpt["config"] model = PCSHODLM(config).to(device) model.load_state_dict(ckpt["model_state_dict"]) model.eval() x_0 = torch.randint(1, config.vocab_size, (2, config.max_seq_len), device=device) # Baseline loss with torch.no_grad(): out_before = model(x_0) loss_before = out_before["loss"].item() energy_before = out_before["energies"][-1] print(f"Before insertion: {model.n_active_layers} layers, loss={loss_before:.4f}, energy={energy_before:.0f}") # Insert 2 layers mgr = DynamicLayerManager(model) mgr.insert_layer(model.n_active_layers // 2, init_strategy="identity") mgr.insert_layer(model.n_active_layers // 2 + 1, init_strategy="identity") print(f"After insertion: {model.n_active_layers} layers, params={count_parameters(model):,}") # Immediate loss (should be similar due to identity init) with torch.no_grad(): out_after = model(x_0) loss_after = out_after["loss"].item() energy_after = out_after["energies"][-1] print(f"Immediate post-insert: loss={loss_after:.4f}, energy={energy_after:.0f}") # Now run online learning to train the new layers print("Running online learning to train new layers...") updater = InferenceUpdater(model, lr=1e-4, grad_clip=0.5) online_losses = [] for step in range(50): t = torch.randint(1, config.n_diffusion_steps + 1, (2,), device=device) x_t, mask = model.schedule.corrupt(x_0, t, config.mask_token_id) h_0 = model.embed_input(x_t, t) h_init = model.amortized_forward_pass(h_0) h_settled, _, energies, _, errors = model.settle( h_init, x_0, mask, t, return_errors=True ) if errors is not None: eps_up, eps_down = errors updater.update_from_errors(h_settled, eps_up, eps_down, x_0, mask, t) with torch.no_grad(): logits = model.readout(model.readout_norm(h_settled[-1])) ml = logits[mask] mt = x_0[mask] if ml.numel() > 0: online_losses.append(F.cross_entropy(ml, mt).item()) if (step + 1) % 10 == 0: print(f" Online step {step+1}: loss={online_losses[-1]:.4f}") loss_final = online_losses[-1] if online_losses else loss_after print(f"After 50 online steps: loss={loss_final:.4f}") results = { "n_layers_before": config.n_layers - 2, "n_layers_after": config.n_layers, "loss_before": loss_before, "loss_immediate_after": loss_after, "loss_after_online": loss_final, "online_loss_trace": online_losses, } return results # ============================================================================= # Main # ============================================================================= def run_all( checkpoint_path: str = None, ao3_data_dir: str = None, output_dir: str = "online_learning_results", device: str = "cpu", ): output_path = Path(output_dir) output_path.mkdir(parents=True, exist_ok=True) # Find checkpoint if checkpoint_path is None: for d in ["checkpoints/local_pc_1k", "checkpoints/local_pc"]: p = os.path.join(d, "checkpoint_final.pt") if os.path.exists(p): checkpoint_path = p break if checkpoint_path is None: print("No checkpoint found. Train a model first.") return # Find AO3 data if ao3_data_dir is None: ao3_data_dir = os.path.join( os.path.dirname(__file__), "..", "data", "ao3-metadata" ) all_results = {} # Experiment 1: Online learning adaptation if os.path.isdir(ao3_data_dir): all_results["online_adaptation"] = test_online_learning_adaptation( checkpoint_path, ao3_data_dir, device ) else: print(f"AO3 data not found at {ao3_data_dir}, skipping experiment 1") # Experiment 2: Dynamic layer insertion all_results["dynamic_insertion"] = test_dynamic_layer_insertion( checkpoint_path, device ) # Save with open(output_path / "results.json", "w") as f: json.dump(all_results, f, indent=2) print(f"\nResults saved to {output_path}/results.json") return all_results if __name__ == "__main__": import argparse parser = argparse.ArgumentParser() parser.add_argument("--checkpoint", type=str, default=None) parser.add_argument("--ao3_data", type=str, default=None) parser.add_argument("--output", type=str, default="online_learning_results") parser.add_argument("--device", type=str, default="cpu") args = parser.parse_args() run_all(args.checkpoint, args.ao3_data, args.output, args.device)