""" 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()