pc-sho-dlm-code / app_optimized.py
Zae
PC-SHO-DLM: full architecture with MSA integration
c2d8a57
Raw History Blame Contribute Delete
9.96 kB
"""
PC-SHO-DLM v2: OPTIMIZED Training on A100
All Tier 1+2 optimizations enabled:
- Spectral preconditioning (per-element precision dynamics)
- Curriculum over K (ramp settling steps)
- Temporal hierarchy (per-layer mass/damping)
- Homeostatic plasticity (adaptive anchoring)
- Linear attention in settling loop
- Lateral inhibition (competitive token settling)
- Adaptive two-timescale ratio (for unified mode)
Trains on massive streaming data from HuggingFace.
"""
import json, os, sys, threading, time, math
import gradio as gr
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset, IterableDataset
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "src"))
from model import PCSHODLM, PCSHOConfig, LocalParameterUpdater, count_parameters
# =============================================================================
# Streaming Dataset — handles terabytes via HF datasets streaming
# =============================================================================
class StreamingCharDataset(IterableDataset):
"""Streams text from HuggingFace datasets, encodes as bytes on the fly."""
def __init__(self, dataset_name, config_name, split, seq_len, vocab_size=257, max_tokens=None):
self.dataset_name = dataset_name
self.config_name = config_name
self.split = split
self.seq_len = seq_len
self.vocab_size = vocab_size
self.max_tokens = max_tokens
def __iter__(self):
from datasets import load_dataset
ds = load_dataset(self.dataset_name, self.config_name, split=self.split, streaming=True)
buffer = []
total = 0
for item in ds:
text = item.get("text", "")
if not text.strip():
continue
encoded = [min(b + 1, self.vocab_size - 1) for b in text.encode("utf-8")]
buffer.extend(encoded)
while len(buffer) >= self.seq_len:
chunk = buffer[:self.seq_len]
buffer = buffer[self.seq_len:]
total += self.seq_len
if self.max_tokens and total > self.max_tokens:
return
yield {"input_ids": torch.tensor(chunk, dtype=torch.long)}
class CharDS(Dataset):
def __init__(self, text, seq_len):
self.seq_len = seq_len
self.data = torch.tensor([min(b+1,256) for b in text.encode("utf-8")], dtype=torch.long)
self.n = max(1, (len(self.data)-seq_len)//seq_len)
def __len__(self): return self.n
def __getitem__(self, i):
return {"input_ids": self.data[i*self.seq_len:(i+1)*self.seq_len]}
# =============================================================================
# Training
# =============================================================================
LOG = []
STATUS = "Idle"
def train():
global STATUS, LOG
LOG = ["PC-SHO-DLM v2 — OPTIMIZED Training"]
device = "cuda" if torch.cuda.is_available() else "cpu"
LOG.append(f"Device: {device}")
if device == "cuda":
LOG.append(f"GPU: {torch.cuda.get_device_name()}")
LOG.append(f"VRAM: {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB")
# 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,
# Temporal hierarchy
mass_scale=0.5, gamma_scale=0.3, eta_scale=0.3,
# Homeostatic
homeostatic_target=1.0, homeostatic_rate=0.001,
# Lateral inhibition
lateral_inhibition=0.01, settling_budget_fraction=0.5,
)
n_params = count_parameters(PCSHODLM(config))
LOG.append(f"Model: {n_params:,} params (d={config.d_model}, L={config.n_layers})")
LOG.append("Optimizations: spectral precond, curriculum-K, temporal hierarchy,")
LOG.append(" homeostatic plasticity, linear-attn settling, lateral inhibition")
# === DATA: Stream from multiple sources ===
STATUS = "Loading data (streaming)..."
LOG.append("\nData sources (streaming):")
# Use WikiText for validation (small, fixed)
from datasets import load_dataset
dv = load_dataset("wikitext", "wikitext-103-raw-v1", split="validation")
val_text = "\n".join([r["text"] for r in dv if r["text"].strip()])
val_ds = CharDS(val_text, config.max_seq_len)
LOG.append(f" Val: WikiText-103 validation ({len(val_text):,} chars)")
# Stream from FineWeb-10BT (10 billion tokens) — ~50GB of high-quality web text
# Falls back to WikiText if FineWeb unavailable
try:
train_ds = StreamingCharDataset(
"HuggingFaceFW/fineweb", "sample-10BT", "train",
seq_len=config.max_seq_len, max_tokens=2_000_000_000 # 2B tokens
)
LOG.append(f" Train: FineWeb-10BT streaming (2B tokens)")
except Exception as e:
LOG.append(f" FineWeb failed ({e}), falling back to WikiText-103")
train_ds = StreamingCharDataset(
"wikitext", "wikitext-103-raw-v1", "train",
seq_len=config.max_seq_len, max_tokens=500_000_000
)
LOG.append(f" Train: WikiText-103 streaming (500M tokens)")
model = PCSHODLM(config).to(device)
model.train()
# Local PC training with all optimizations
updater = LocalParameterUpdater(model, lr_forward=3e-4, lr_feedback=3e-4, lr_readout=3e-4, lr_precision=3e-5)
tl = DataLoader(train_ds, batch_size=64, num_workers=4, pin_memory=True)
vl = DataLoader(val_ds, batch_size=64, shuffle=False, num_workers=2, pin_memory=True)
STATUS = "Training OPTIMIZED Local PC..."
LOG.append(f"\n{'='*60}")
LOG.append(f"Training: OPTIMIZED Local PC | {n_params:,} params")
LOG.append(f"{'='*60}")
step, start, max_steps = 0, time.time(), 10000
log_data = {"steps": [], "losses": [], "energies": [], "K_values": []}
K_max = config.n_settling_steps
for batch in tl:
if step >= max_steps:
break
x_0 = batch["input_ids"].to(device)
# CURRICULUM OVER K
K_warmup = min(2000, max_steps // 3)
K_curr = max(1, int(K_max * min(1.0, step / K_warmup)))
model.config.n_settling_steps = K_curr
result = updater.step({"input_ids": x_0})
step += 1
loss = result.get("loss", 0)
energies = result.get("energies", [])
if step % 100 == 0:
elapsed = time.time() - start
tps = step * 64 * config.max_seq_len / elapsed
e = f"{energies[-1]:.0f}" if energies else "N/A"
rho_str = f"[{model.layer_rho[0]:.4f}..{model.layer_rho[-1]:.4f}]"
msg = f"[opt-pc] Step {step:6d} | Loss: {loss:.4f} | Energy: {e} | K={K_curr} | rho={rho_str} | Tok/s: {tps:.0f} | {elapsed:.0f}s"
LOG.append(msg)
log_data["steps"].append(step)
log_data["losses"].append(loss)
log_data["energies"].append(energies[-1] if energies else 0)
log_data["K_values"].append(K_curr)
if step % 2000 == 0:
model.eval()
tl2, tt2 = 0.0, 0
with torch.no_grad():
for vb in vl:
vx = vb["input_ids"].to(device)
vo = model(vx)
nm = vo["mask"].sum().item()
if nm > 0:
tl2 += vo["loss"].item() * nm
tt2 += nm
if tt2 > 100000:
break
vl2 = tl2 / max(1, tt2)
LOG.append(f" --> Val loss: {vl2:.4f}")
log_data.setdefault("val", []).append((step, vl2))
model.train()
# Restore full K for final eval
model.config.n_settling_steps = K_max
model.eval()
tl2, tt2 = 0.0, 0
with torch.no_grad():
for vb in vl:
vx = vb["input_ids"].to(device)
vo = model(vx)
nm = vo["mask"].sum().item()
if nm > 0:
tl2 += vo["loss"].item() * nm
tt2 += nm
if tt2 > 200000:
break
fv = tl2 / max(1, tt2)
log_data["final_val"] = fv
elapsed = time.time() - start
LOG.append(f"\nDONE: {step} steps in {elapsed:.0f}s | Final val: {fv:.4f}")
STATUS = "Complete!"
os.makedirs("results", exist_ok=True)
torch.save({"model": model.state_dict(), "config": config, "log": log_data}, "results/optimized_pc_10k.pt")
with open("results/optimized_log.json", "w") as f:
json.dump(log_data, f)
try:
from huggingface_hub import HfApi
HfApi().upload_folder(folder_path="results", repo_id="zotowata/pc-sho-dlm-optimized",
repo_type="space", path_in_repo="results")
LOG.append("Results uploaded!")
except Exception as e:
LOG.append(f"Upload: {e}")
_thread = None
def start():
global _thread
if _thread and _thread.is_alive():
return "Already running!"
_thread = threading.Thread(target=train, daemon=True)
_thread.start()
return "OPTIMIZED training started on A100!"
with gr.Blocks(title="PC-SHO-DLM v2 Optimized") as demo:
gr.Markdown("# PC-SHO-DLM v2: OPTIMIZED Local PC Training (A100)")
gr.Markdown("All Tier 1+2 optimizations: spectral precond, curriculum-K, temporal hierarchy, homeostasis, linear-attn settling, lateral inhibition")
with gr.Row():
btn = gr.Button("Start Training", variant="primary")
st = gr.Textbox(label="Status", value="Idle")
log = gr.Textbox(label="Log", lines=30, max_lines=60)
btn.click(start, outputs=st)
refresh = gr.Button("Refresh")
refresh.click(lambda: "\n".join(LOG[-60:]), outputs=log)
refresh.click(lambda: STATUS, outputs=st)
timer = gr.Timer(5)
timer.tick(lambda: "\n".join(LOG[-60:]), outputs=log)
timer.tick(lambda: STATUS, outputs=st)
demo.launch()