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