pc-sho-dlm-code / app.py
Zae
PC-SHO-DLM: full architecture with MSA integration
c2d8a57
Raw History Blame Contribute Delete
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()