Download app.py from zotowata/pc-sho-dlm-code: direct link, hf CLI and curl.
- Browser
- Download file 10 kB
-
https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/app.py
- Command line
-
hf download hf://zotowata/pc-sho-dlm-code/app.py
-
curl -L -o app.py https://huggingface.co/zotowata/pc-sho-dlm-code/resolve/main/app.py
10 kB
| """ | |
| PC-SHO-DLM A100 Training — HuggingFace Spaces | |
| Runs all 3 training modes and displays results via Gradio. | |
| """ | |
| import json | |
| import os | |
| import sys | |
| import threading | |
| import time | |
| import gradio as gr | |
| import torch | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, Dataset | |
| # Import model | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "src")) | |
| from model import PCSHODLM, PCSHOConfig, LocalParameterUpdater, count_parameters | |
| # ============================================================================= | |
| # Dataset | |
| # ============================================================================= | |
| class CharLevelDataset(Dataset): | |
| def __init__(self, text, seq_len, vocab_size=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]} | |
| # ============================================================================= | |
| # Training | |
| # ============================================================================= | |
| TRAINING_LOG = [] | |
| CURRENT_STATUS = "Idle" | |
| def train_one_mode(mode, config, train_ds, val_ds, max_steps=10000, batch_size=64, lr=3e-4): | |
| global CURRENT_STATUS, TRAINING_LOG | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| CURRENT_STATUS = f"Training {mode} on {device}..." | |
| TRAINING_LOG.append(f"\n{'='*60}") | |
| TRAINING_LOG.append(f"Mode: {mode} | Params: {count_parameters(PCSHODLM(config)):,} | Device: {device}") | |
| TRAINING_LOG.append(f"{'='*60}") | |
| model = PCSHODLM(config).to(device) | |
| model.train() | |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, drop_last=True, | |
| num_workers=4 if device == "cuda" else 0, pin_memory=device == "cuda") | |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, | |
| num_workers=2 if device == "cuda" else 0, pin_memory=device == "cuda") | |
| 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) | |
| log = {"mode": mode, "steps": [], "losses": [], "energies": [], "val_losses": []} | |
| step = 0 | |
| start = time.time() | |
| 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() | |
| out = model(x_0) | |
| loss = out["loss"] | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| scheduler.step() | |
| loss_val = loss.item() | |
| energies = out["energies"] | |
| elif mode == "local": | |
| result = updater.step({"input_ids": x_0}) | |
| 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_s, _, energies = model.unified_settle(h_init, x_0, mask, t, param_lr_scale=0.005) | |
| with torch.no_grad(): | |
| logits = model.readout(model.readout_norm(h_s[-1])) | |
| ml, mt = logits[mask], x_0[mask] | |
| loss_val = F.cross_entropy(ml, mt).item() if ml.numel() > 0 else 0 | |
| step += 1 | |
| if step % 200 == 0: | |
| elapsed = time.time() - start | |
| tps = (step * batch_size * config.max_seq_len) / elapsed | |
| energy_str = f"{energies[-1]:.0f}" if energies else "N/A" | |
| msg = f"[{mode}] Step {step:6d} | Loss: {loss_val:.4f} | Energy: {energy_str} | Tok/s: {tps:.0f} | {elapsed:.0f}s" | |
| TRAINING_LOG.append(msg) | |
| log["steps"].append(step) | |
| log["losses"].append(loss_val) | |
| log["energies"].append(energies[-1] if energies else 0) | |
| if step % 2000 == 0: | |
| model.eval() | |
| tl, tt = 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: | |
| tl += vo["loss"].item() * nm | |
| tt += nm | |
| if tt > 100000: | |
| break | |
| vl = tl / max(1, tt) | |
| TRAINING_LOG.append(f" --> [{mode}] Val loss: {vl:.4f}") | |
| log["val_losses"].append((step, vl)) | |
| model.train() | |
| elapsed = time.time() - start | |
| # Final val | |
| model.eval() | |
| tl, tt = 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: | |
| tl += vo["loss"].item() * nm | |
| tt += nm | |
| if tt > 200000: | |
| break | |
| final_val = tl / max(1, tt) | |
| log["final_val_loss"] = final_val | |
| msg = f"[{mode}] DONE: {step} steps in {elapsed:.0f}s | Final val loss: {final_val:.4f}" | |
| TRAINING_LOG.append(msg) | |
| # Save checkpoint | |
| os.makedirs("results", exist_ok=True) | |
| torch.save({"model": model.state_dict(), "config": config, "log": log}, f"results/{mode}_10k.pt") | |
| return log | |
| def run_all_training(): | |
| global CURRENT_STATUS, TRAINING_LOG | |
| TRAINING_LOG = ["Starting PC-SHO-DLM A100 Training..."] | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| TRAINING_LOG.append(f"Device: {device}") | |
| if device == "cuda": | |
| TRAINING_LOG.append(f"GPU: {torch.cuda.get_device_name()}") | |
| TRAINING_LOG.append(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") | |
| # A100-optimized config | |
| 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, | |
| ) | |
| TRAINING_LOG.append(f"Model: d={config.d_model}, L={config.n_layers}, ~{count_parameters(PCSHODLM(config)):,} params") | |
| # Load data | |
| CURRENT_STATUS = "Loading WikiText-103..." | |
| from datasets import load_dataset | |
| 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()])[:100_000_000] | |
| val_text = "\n".join([r["text"] for r in ds_val if r["text"].strip()]) | |
| train_ds = CharLevelDataset(train_text, config.max_seq_len) | |
| val_ds = CharLevelDataset(val_text, config.max_seq_len) | |
| TRAINING_LOG.append(f"Data: {len(train_text):,} train chars, {len(val_text):,} val chars") | |
| results = {} | |
| for mode in ["backprop", "local", "unified"]: | |
| results[mode] = train_one_mode(mode, config, train_ds, val_ds, max_steps=10000, batch_size=64) | |
| # Final summary | |
| TRAINING_LOG.append(f"\n{'='*60}") | |
| TRAINING_LOG.append("FINAL RESULTS") | |
| TRAINING_LOG.append(f"{'='*60}") | |
| for mode, log in results.items(): | |
| TRAINING_LOG.append(f"{mode:15s} | Val loss: {log.get('final_val_loss', 'N/A')}") | |
| CURRENT_STATUS = "Complete!" | |
| # Upload results | |
| try: | |
| from huggingface_hub import HfApi | |
| api = HfApi() | |
| with open("results/results.json", "w") as f: | |
| json.dump(results, f, indent=2) | |
| api.upload_folder(folder_path="results", repo_id="zotowata/pc-sho-dlm-train", | |
| repo_type="space", path_in_repo="results") | |
| TRAINING_LOG.append("Results uploaded to HuggingFace!") | |
| except Exception as e: | |
| TRAINING_LOG.append(f"Upload error: {e}") | |
| # ============================================================================= | |
| # Gradio UI | |
| # ============================================================================= | |
| training_thread = None | |
| def start_training(): | |
| global training_thread | |
| if training_thread and training_thread.is_alive(): | |
| return "Training already running!" | |
| training_thread = threading.Thread(target=run_all_training, daemon=True) | |
| training_thread.start() | |
| return "Training started on A100!" | |
| def get_log(): | |
| return "\n".join(TRAINING_LOG[-50:]) | |
| def get_status(): | |
| return CURRENT_STATUS | |
| with gr.Blocks(title="PC-SHO-DLM Training") as demo: | |
| gr.Markdown("# PC-SHO-DLM: Predictive-Coding Diffusion LM — A100 Training") | |
| gr.Markdown("Trains 3 modes (backprop, local PC, unified) on WikiText-103 with ~50M params") | |
| with gr.Row(): | |
| start_btn = gr.Button("Start Training", variant="primary") | |
| status = gr.Textbox(label="Status", value="Idle") | |
| log_box = gr.Textbox(label="Training Log", lines=25, max_lines=50) | |
| refresh_btn = gr.Button("Refresh Log") | |
| start_btn.click(start_training, outputs=status) | |
| refresh_btn.click(get_log, outputs=log_box) | |
| refresh_btn.click(get_status, outputs=status) | |
| # Auto-refresh via Timer | |
| timer = gr.Timer(5) | |
| timer.tick(get_log, outputs=log_box) | |
| timer.tick(get_status, outputs=status) | |
| demo.launch() | |