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