""" PC-SHO-DLM GPU Training Script (A100 optimized) Trains a scaled-up PC-SHO-DLM on WikiText-103 using CUDA. Designed for HuggingFace Spaces with A100 (80GB) GPU. Runs all 3 training modes sequentially and saves results. """ import json import os import sys import time import torch import torch.nn.functional as F from torch.utils.data import DataLoader, Dataset sys.path.insert(0, "/app/src") from model import PCSHODLM, PCSHOConfig, LocalParameterUpdater, count_parameters # ============================================================================= # Dataset # ============================================================================= class CharLevelDataset(Dataset): def __init__(self, text: str, seq_len: int, vocab_size: int = 257): self.seq_len = seq_len 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 return {"input_ids": self.data[start : start + self.seq_len]} def load_data(seq_len=512, max_chars=100_000_000): """Load WikiText-103 from HuggingFace.""" from datasets import load_dataset print("Loading WikiText-103...") ds_train = load_dataset("wikitext", "wikitext-103-raw-v1", split="train") ds_val = load_dataset("wikitext", "wikitext-103-raw-v1", split="validation") train_text = "\n".join([r["text"] for r in ds_train if r["text"].strip()])[:max_chars] val_text = "\n".join([r["text"] for r in ds_val if r["text"].strip()]) print(f"Train: {len(train_text):,} chars, Val: {len(val_text):,} chars") return CharLevelDataset(train_text, seq_len), CharLevelDataset(val_text, seq_len) # ============================================================================= # Training # ============================================================================= def train_model(model, train_ds, val_ds, mode, config, max_steps=10000, batch_size=64, lr=3e-4): device = "cuda" model = model.to(device) model.train() train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, drop_last=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True) if mode == "backprop": optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_steps, eta_min=lr * 0.1) elif mode == "local": updater = LocalParameterUpdater(model, lr_forward=lr, lr_feedback=lr, lr_readout=lr, lr_precision=lr * 0.1) elif mode == "unified": param_lr_scale = 0.005 log = {"step": [], "loss": [], "energy": [], "val_loss": [], "wall_time": []} step = 0 start = time.time() print(f"\n{'='*70}") print(f"Training PC-SHO-DLM | Mode: {mode} | Params: {count_parameters(model):,}") print(f"Device: {torch.cuda.get_device_name()} | Batch: {batch_size} | Steps: {max_steps}") print(f"{'='*70}") while step < max_steps: for batch in train_loader: if step >= max_steps: break x_0 = batch["input_ids"].to(device) if mode == "backprop": optimizer.zero_grad() output = model(x_0) loss = output["loss"] loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() loss_val = loss.item() energies = output["energies"] elif mode == "local": batch_d = {"input_ids": x_0} result = updater.step(batch_d) loss_val = result.get("loss", 0.0) energies = result.get("energies", []) elif mode == "unified": B, S = x_0.shape t = torch.randint(1, config.n_diffusion_steps + 1, (B,), 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 = model.unified_settle(h_init, x_0, mask, t, param_lr_scale=param_lr_scale) with torch.no_grad(): logits = model.readout(model.readout_norm(h_settled[-1])) ml, mt = logits[mask], x_0[mask] loss_val = F.cross_entropy(ml, mt).item() if ml.numel() > 0 else 0.0 step += 1 if step % 100 == 0: elapsed = time.time() - start energy_str = f"{energies[-1]:.0f}" if energies else "N/A" tps = (step * batch_size * config.max_seq_len) / elapsed print(f"Step {step:6d} | Loss: {loss_val:.4f} | Energy: {energy_str} | Tok/s: {tps:.0f} | Time: {elapsed:.0f}s") log["step"].append(step) log["loss"].append(loss_val) log["energy"].append(energies[-1] if energies else 0) log["wall_time"].append(elapsed) if step % 2000 == 0: # Validation model.eval() total_loss, total_tok = 0.0, 0 with torch.no_grad(): for vb in val_loader: vx = vb["input_ids"].to(device) vo = model(vx) nm = vo["mask"].sum().item() if nm > 0: total_loss += vo["loss"].item() * nm total_tok += nm if total_tok > 100000: break val_loss = total_loss / max(1, total_tok) print(f" --> Val loss: {val_loss:.4f}") log["val_loss"].append((step, val_loss)) model.train() elapsed = time.time() - start print(f"Done. {step} steps in {elapsed:.0f}s ({step*batch_size*config.max_seq_len/elapsed:.0f} tok/s)") # Final validation model.eval() total_loss, total_tok = 0.0, 0 with torch.no_grad(): for vb in val_loader: vx = vb["input_ids"].to(device) vo = model(vx) nm = vo["mask"].sum().item() if nm > 0: total_loss += vo["loss"].item() * nm total_tok += nm if total_tok > 200000: break final_val = total_loss / max(1, total_tok) print(f"Final val loss: {final_val:.4f}") log["final_val_loss"] = final_val return model, log # ============================================================================= # Main # ============================================================================= def main(): torch.backends.cudnn.benchmark = True # A100-optimized config: bigger model, bigger batch, longer sequences config = PCSHOConfig( vocab_size=257, max_seq_len=512, d_model=512, n_heads=8, n_layers=12, d_ff=2048, n_diffusion_steps=128, n_settling_steps=6, mask_token_id=0, dropout=0.1, feedback_rank=128, ) # Load data train_ds, val_ds = load_data(seq_len=config.max_seq_len, max_chars=100_000_000) results = {} save_dir = "/app/results" os.makedirs(save_dir, exist_ok=True) # 1. Backprop baseline print("\n" + "=" * 70) print("PHASE 1: Backprop Baseline") print("=" * 70) model_bp = PCSHODLM(config) model_bp, log_bp = train_model(model_bp, train_ds, val_ds, "backprop", config, max_steps=10000, batch_size=64, lr=3e-4) results["backprop"] = log_bp torch.save({"model": model_bp.state_dict(), "config": config, "log": log_bp}, f"{save_dir}/backprop_10k.pt") # 2. Local PC print("\n" + "=" * 70) print("PHASE 2: Local PC (globally backprop-free)") print("=" * 70) model_pc = PCSHODLM(config) model_pc, log_pc = train_model(model_pc, train_ds, val_ds, "local", config, max_steps=10000, batch_size=64, lr=3e-4) results["local_pc"] = log_pc torch.save({"model": model_pc.state_dict(), "config": config, "log": log_pc}, f"{save_dir}/local_pc_10k.pt") # 3. Unified (settling=learning) print("\n" + "=" * 70) print("PHASE 3: Unified (settling = learning)") print("=" * 70) model_uni = PCSHODLM(config) model_uni, log_uni = train_model(model_uni, train_ds, val_ds, "unified", config, max_steps=10000, batch_size=64, lr=3e-4) results["unified"] = log_uni torch.save({"model": model_uni.state_dict(), "config": config, "log": log_uni}, f"{save_dir}/unified_10k.pt") # Save combined results with open(f"{save_dir}/results.json", "w") as f: json.dump(results, f, indent=2) # Print final comparison print("\n" + "=" * 70) print("FINAL RESULTS") print("=" * 70) for mode, log in results.items(): final_loss = log.get("final_val_loss", log["loss"][-1] if log["loss"] else "N/A") print(f"{mode:15s} | Final val loss: {final_loss}") # Push results to HF try: from huggingface_hub import HfApi api = HfApi() api.upload_folder( folder_path=save_dir, repo_id="zotowata/pc-sho-dlm-train", repo_type="space", path_in_repo="results", ) print("\nResults uploaded to HuggingFace!") except Exception as e: print(f"Upload failed: {e}") if __name__ == "__main__": main()