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