File size: 9,959 Bytes
c2d8a57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
"""
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()