pc-sho-dlm-code / experiments /online_learning_eval.py
Zae
PC-SHO-DLM: full architecture with MSA integration
c2d8a57
Raw History Blame Contribute Delete
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)