""" PC-SHO-DLM Training Script Supports two training modes: 1. Standard (backprop) training - for baselines and comparison 2. Local (PC) training - the proposed globally backprop-free method Supports data sources: - HuggingFace datasets (wikitext, etc.) - Raw text files - Synthetic data (for testing) """ import argparse import json import math import os import time from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset from model import PCSHODLM, PCSHOConfig, LocalParameterUpdater, count_parameters # ============================================================================= # Datasets # ============================================================================= class CharLevelDataset(Dataset): """Character-level dataset for Stage 1 proof-of-mechanism experiments.""" def __init__(self, text: str, seq_len: int, vocab_size: int = 256): self.seq_len = seq_len self.vocab_size = vocab_size # Encode as bytes, offset by 1 (0 = MASK token) self.data = torch.tensor( [min(b + 1, vocab_size - 1) for b in text.encode("utf-8")], dtype=torch.long, ) self.n_seqs = max(1, (len(self.data) - seq_len) // seq_len) def __len__(self): return self.n_seqs def __getitem__(self, idx): start = idx * self.seq_len end = start + self.seq_len return {"input_ids": self.data[start:end]} def load_wikitext(seq_len: int, vocab_size: int = 257, split: str = "train"): """Load WikiText-103 from HuggingFace.""" from datasets import load_dataset print(f"Loading WikiText-103 ({split})...") ds = load_dataset("wikitext", "wikitext-103-raw-v1", split=split) # Concatenate all text text = "\n".join([row["text"] for row in ds if row["text"].strip()]) print(f" {len(text):,} characters loaded") return CharLevelDataset(text, seq_len, vocab_size) def load_text_file(path: str, seq_len: int, vocab_size: int = 257, max_chars: int = 0): """Load from a raw text file. Args: max_chars: limit characters loaded (0 = all). Use for large files. """ with open(path, "r", errors="replace") as f: if max_chars > 0: text = f.read(max_chars) else: text = f.read() print(f"Loaded {len(text):,} characters from {path}") return CharLevelDataset(text, seq_len, vocab_size) # ============================================================================= # Training Loop # ============================================================================= class Trainer: """Trainer supporting both standard backprop and local PC training. Tracks all metrics recommended by the paper's diagnostics section: - Loss (masked token NLL) - Energy trace (per-step latent energy) - Stationarity residual (envelope theorem validation) - Energy monotonicity (Lyapunov check) """ def __init__( self, model: PCSHODLM, train_dataset: Dataset, val_dataset: Dataset = None, training_mode: str = "local", lr: float = 1e-4, batch_size: int = 16, max_steps: int = 10000, log_interval: int = 50, eval_interval: int = 500, save_interval: int = 2000, save_dir: str = "checkpoints", device: str = "cpu", grad_clip: float = 1.0, ): self.model = model.to(device) self.device = device self.training_mode = training_mode self.max_steps = max_steps self.log_interval = log_interval self.eval_interval = eval_interval self.save_interval = save_interval self.save_dir = Path(save_dir) self.save_dir.mkdir(parents=True, exist_ok=True) self.grad_clip = grad_clip self.train_loader = DataLoader( train_dataset, batch_size=batch_size, shuffle=True, drop_last=True, num_workers=0, ) self.val_loader = ( DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=0) if val_dataset else None ) if training_mode == "local": self.updater = LocalParameterUpdater( model, lr_forward=lr, lr_feedback=lr, lr_readout=lr, lr_precision=lr * 0.1, ) elif training_mode == "unified": self.param_lr = lr else: self.optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01) self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( self.optimizer, T_max=max_steps, eta_min=lr * 0.1 ) # Logging self.log = { "step": [], "loss": [], "energy_initial": [], "energy_final": [], "stationarity_residual": [], "wall_time": [], "val_loss": [], "val_loss_settled": [], "amortized_loss": [], "tokens_per_sec": [], } self.running_loss = 0.0 self.running_amortized_loss = 0.0 self.running_energy_init = 0.0 self.running_energy_final = 0.0 self.running_stationarity = 0.0 self.running_count = 0 def train_step_backprop(self, batch): """Standard end-to-end backprop training step.""" self.optimizer.zero_grad() x_0 = batch["input_ids"].to(self.device) output = self.model(x_0) loss = output["loss"] loss.backward() torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.grad_clip) self.optimizer.step() self.scheduler.step() return { "loss": loss.item(), "energies": output["energies"], "stationarity": output.get("stationarity_residual", 0.0), } def train_step_local(self, batch): """Local predictive-coding training step (globally backprop-free).""" batch = {k: v.to(self.device) for k, v in batch.items()} result = self.updater.step(batch) return result def train_step_unified(self, batch): """Unified training: settle first, then update all trainable components.""" x_0 = batch["input_ids"].to(self.device) return self.model.unified_train_batch(x_0, param_lr=self.param_lr) @torch.no_grad() def evaluate(self, settled: bool = False): """Evaluate on validation set.""" if self.val_loader is None: return None self.model.eval() total_loss = 0.0 total_tokens = 0 for batch in self.val_loader: x_0 = batch["input_ids"].to(self.device) output = self.model.settled_forward(x_0) if settled else self.model(x_0) n_masked = output["mask"].sum().item() if n_masked > 0: total_loss += output["loss"].item() * n_masked total_tokens += n_masked self.model.train() return total_loss / max(1, total_tokens) def train(self): """Main training loop.""" self.model.train() step = 0 start_time = time.time() epoch = 0 tokens_processed = 0 print(f"{'='*60}") print(f"PC-SHO-DLM Training ({self.training_mode} mode)") print(f"{'='*60}") print(f"Parameters: {count_parameters(self.model):,}") print(f"Device: {self.device}") print(f"Max steps: {self.max_steps}") print(f"Settling steps (K): {self.model.config.n_settling_steps}") print(f"Diffusion steps (T): {self.model.config.n_diffusion_steps}") print(f"{'='*60}") while step < self.max_steps: epoch += 1 for batch in self.train_loader: if step >= self.max_steps: break batch_tokens = batch["input_ids"].numel() # CURRICULUM OVER K: ramp settling steps, minimum K=3 K_max = self.model.config.n_settling_steps K_warmup = min(1000, self.max_steps // 5) K_curr = max(3, int(K_max * min(1.0, step / K_warmup))) self.model.config.n_settling_steps = K_curr if self.training_mode == "backprop": result = self.train_step_backprop(batch) elif self.training_mode == "unified": result = self.train_step_unified(batch) else: result = self.train_step_local(batch) step += 1 tokens_processed += batch_tokens # Accumulate metrics loss_val = result.get("loss", 0.0) amortized_loss_val = result.get("amortized_loss", loss_val) energies = result.get("energies", []) stationarity = result.get("stationarity", 0.0) if isinstance(loss_val, (int, float)): self.running_loss += loss_val if isinstance(amortized_loss_val, (int, float)): self.running_amortized_loss += amortized_loss_val if energies: self.running_energy_init += energies[0] self.running_energy_final += energies[-1] if stationarity: self.running_stationarity += stationarity self.running_count += 1 # Logging if step % self.log_interval == 0 and self.running_count > 0: elapsed = time.time() - start_time avg_loss = self.running_loss / self.running_count avg_amortized = self.running_amortized_loss / self.running_count avg_e_init = self.running_energy_init / self.running_count avg_e_final = self.running_energy_final / self.running_count avg_stat = self.running_stationarity / self.running_count tps = tokens_processed / max(elapsed, 1e-6) energy_reduction = (1 - avg_e_final / max(avg_e_init, 1e-6)) * 100 msg = ( f"Step {step:6d} | " f"Loss: {avg_loss:.4f} | " f"Energy: {avg_e_final:.0f} ({energy_reduction:+.1f}%) | " f"Stat: {avg_stat:.1f} | " f"Tok/s: {tps:.0f} | " f"Time: {elapsed:.0f}s" ) if self.training_mode == "unified": msg = ( f"Step {step:6d} | " f"Settled: {avg_loss:.4f} | " f"Amortized: {avg_amortized:.4f} | " f"Energy: {avg_e_final:.0f} ({energy_reduction:+.1f}%) | " f"Stat: {avg_stat:.1f} | " f"Tok/s: {tps:.0f} | " f"Time: {elapsed:.0f}s" ) print(msg) self.log["step"].append(step) self.log["loss"].append(avg_loss) self.log["amortized_loss"].append(avg_amortized) self.log["energy_initial"].append(avg_e_init) self.log["energy_final"].append(avg_e_final) self.log["stationarity_residual"].append(avg_stat) self.log["wall_time"].append(elapsed) self.log["tokens_per_sec"].append(tps) # Reset running averages self.running_loss = 0.0 self.running_amortized_loss = 0.0 self.running_energy_init = 0.0 self.running_energy_final = 0.0 self.running_stationarity = 0.0 self.running_count = 0 # Evaluation if step % self.eval_interval == 0: val_loss = self.evaluate() if val_loss is not None: print(f" --> Val loss: {val_loss:.4f}") self.log["val_loss"].append((step, val_loss)) if self.training_mode in {"local", "unified"}: val_loss_settled = self.evaluate(settled=True) if val_loss_settled is not None: print(f" --> Val settled loss: {val_loss_settled:.4f}") self.log["val_loss_settled"].append((step, val_loss_settled)) # Save if step % self.save_interval == 0: self.save_checkpoint(step) # Final save self.save_checkpoint(step, final=True) self.save_log() elapsed = time.time() - start_time print(f"{'='*60}") print(f"Training complete. {step} steps in {elapsed:.0f}s") print(f"Final avg loss: {self.log['loss'][-1]:.4f}" if self.log['loss'] else "") print(f"{'='*60}") def save_checkpoint(self, step, final=False): name = "final" if final else f"step_{step}" path = self.save_dir / f"checkpoint_{name}.pt" torch.save( { "step": step, "model_state_dict": self.model.state_dict(), "config": self.model.config, "training_mode": self.training_mode, }, path, ) def save_log(self): path = self.save_dir / "training_log.json" with open(path, "w") as f: json.dump(self.log, f, indent=2) # ============================================================================= # Main # ============================================================================= def main(): parser = argparse.ArgumentParser(description="Train PC-SHO-DLM") parser.add_argument("--mode", choices=["local", "backprop", "unified"], default="local", help="Training mode: 'local' (PC), 'backprop' (baseline), or 'unified' (settle-then-update)") parser.add_argument("--data", type=str, default="wikitext", help="Data source: 'wikitext', 'synthetic', or path to text file") parser.add_argument("--d_model", type=int, default=256) parser.add_argument("--n_layers", type=int, default=6) parser.add_argument("--n_heads", type=int, default=8) parser.add_argument("--seq_len", type=int, default=256) parser.add_argument("--batch_size", type=int, default=16) parser.add_argument("--lr", type=float, default=3e-4) parser.add_argument("--max_steps", type=int, default=10000) parser.add_argument("--n_settling", type=int, default=6) parser.add_argument("--n_diffusion", type=int, default=100) parser.add_argument("--save_dir", type=str, default="checkpoints") parser.add_argument("--device", type=str, default="auto") parser.add_argument("--log_interval", type=int, default=50) parser.add_argument("--eval_interval", type=int, default=500) args = parser.parse_args() # Auto-detect device if args.device == "auto": if torch.cuda.is_available(): args.device = "cuda" elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): args.device = "mps" else: args.device = "cpu" # Config config = PCSHOConfig( vocab_size=257, # 256 bytes + 1 MASK token max_seq_len=args.seq_len, d_model=args.d_model, n_heads=args.n_heads, n_layers=args.n_layers, d_ff=args.d_model * 4, n_diffusion_steps=args.n_diffusion, n_settling_steps=args.n_settling, mask_token_id=0, dropout=0.1, ) # Data wikitext_dir = os.path.join(os.path.dirname(__file__), "..", "data", "wikitext-103") if args.data == "wikitext" and os.path.isdir(wikitext_dir): # Use local wikitext-103 files (first 50M chars for train to fit memory) train_dataset = load_text_file( os.path.join(wikitext_dir, "wiki.train.tokens"), args.seq_len, config.vocab_size, max_chars=50_000_000 ) val_dataset = load_text_file( os.path.join(wikitext_dir, "wiki.valid.tokens"), args.seq_len, config.vocab_size ) elif args.data == "wikitext": train_dataset = load_wikitext(args.seq_len, config.vocab_size, split="train") val_dataset = load_wikitext(args.seq_len, config.vocab_size, split="validation") elif args.data == "synthetic": print("Using synthetic data for testing.") text = "The quick brown fox jumps over the lazy dog. " * 5000 full = CharLevelDataset(text, args.seq_len, config.vocab_size) n_val = max(1, len(full) // 10) train_dataset, val_dataset = torch.utils.data.random_split( full, [len(full) - n_val, n_val] ) elif os.path.exists(args.data): full = load_text_file(args.data, args.seq_len, config.vocab_size) n_val = max(1, len(full) // 10) train_dataset, val_dataset = torch.utils.data.random_split( full, [len(full) - n_val, n_val] ) else: raise ValueError(f"Unknown data source: {args.data}") # Model model = PCSHODLM(config) print(f"\nModel: PC-SHO-DLM ({args.mode} training)") print(f"Parameters: {count_parameters(model):,}") print(f"Architecture: d={config.d_model}, L={config.n_layers}, heads={config.n_heads}") print(f"Settling: K={config.n_settling_steps}, T={config.n_diffusion_steps}") print(f"Sequence length: {config.max_seq_len}") print(f"Train samples: {len(train_dataset):,}") # Train trainer = Trainer( model=model, train_dataset=train_dataset, val_dataset=val_dataset, training_mode=args.mode, lr=args.lr, batch_size=args.batch_size, max_steps=args.max_steps, log_interval=args.log_interval, eval_interval=args.eval_interval, save_dir=args.save_dir, device=args.device, ) trainer.train() if __name__ == "__main__": main()