pc-sho-dlm-code / src /train.py
zotowata's picture
Sync 2B training job files
2a5274f verified
Raw History Blame Contribute Delete
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)
@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()