Download experiments/online_learning_eval.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/experiments/online_learning_eval.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/experiments/online_learning_eval.py
-
curl -L -o online_learning_eval.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/experiments/online_learning_eval.py
11.8 kB
| """ | |
| 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) | |