AbstractPhil commited on
Commit
b709283
·
verified ·
1 Parent(s): e38d143

Create notebook.py

Browse files
Files changed (1) hide show
  1. notebook.py +1540 -0
notebook.py ADDED
@@ -0,0 +1,1540 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Pentachoron Constellation — Multi-Channel, HF Push, and Dataset Sweep
4
+ V2
5
+ Apache-2.0
6
+ Author: AbstractPhil
7
+ Quartermaster: Mirel (GPT-5 Thinking) + Claude Sonnet 4.5
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import os, sys, json, math, time, random, shutil, zipfile, platform
13
+ from pathlib import Path
14
+ from datetime import datetime
15
+ from typing import List, Tuple, Dict, Optional
16
+
17
+ import numpy as np
18
+ import torch
19
+ import torch.nn as nn
20
+ import torch.nn.functional as F
21
+ from torchvision import datasets, transforms
22
+ from torch.utils.data import DataLoader
23
+ from tqdm import tqdm
24
+ from sklearn.metrics import confusion_matrix
25
+
26
+ ## ---------------------------------------------------------------------
27
+ ## Fast settings / safety
28
+ ## ---------------------------------------------------------------------
29
+ #torch.autograd.set_detect_anomaly(False)
30
+ #if torch.cuda.is_available():
31
+ # torch.backends.cudnn.benchmark = True
32
+ # torch.cuda.empty_cache()
33
+
34
+
35
+ # ---------------------------------------------------------------------
36
+ # Configuration (edit these)
37
+ # ---------------------------------------------------------------------
38
+ config: Dict = {
39
+ # Model dims
40
+ "input_dim": 28*28, # will be set by loader
41
+ "input_channels": "auto", # "auto" | 1 | 3 ; loader enforces
42
+ "base_dim": 28*25,
43
+ "proj_dim": None,
44
+
45
+ # Constellation
46
+ "num_classes": 10, # set by loader
47
+ "num_pentachoron_pairs": 4,
48
+ "lambda_separation": 0.391,
49
+
50
+ # Attention / extractor
51
+ "num_heads": 4,
52
+ "channels": 24,
53
+
54
+ # Training
55
+ "batch_size": 1024,
56
+ "epochs": 20,
57
+ "lr": 3e-3,
58
+ "weight_decay": 1e-5,
59
+ "temp": 0.5,
60
+
61
+ # Loss weights
62
+ "w_ce": 1.0,
63
+ "w_dual": 1.0,
64
+ "w_rose": 1.0,
65
+ "w_diag": 0.1,
66
+ "w_reg": 0.1,
67
+
68
+ # Legacy compat
69
+ "loss_weight_scalar": 1.0,
70
+
71
+ # Dataset override knobs
72
+ "img_size": 28, # unified target size
73
+ "img_channels": "auto", # "auto" | 1 | 3 ; coerces all sets
74
+ "normalize": True,
75
+ "per_dataset_norm": True,
76
+ "augment": True, # safe light aug
77
+
78
+ # Sweep control
79
+ "sweep_all": False, # set True to run all datasets
80
+ "seed": 420,
81
+
82
+ # Hugging Face
83
+ "hf_repo_id": "AbstractPhil/pentachora-multi-channel-frequency-encoded-2",
84
+ "dataset": "FashionMNIST",
85
+ }
86
+
87
+ # --- HF pathing / naming ---
88
+ config.setdefault("hf_subdir_root", "")
89
+ config.setdefault("hf_dataset_dir_template", "{dataset}") # folder under root
90
+ config.setdefault("hf_run_dir_template", "{ts}_{dataset}") # or "{ts}"
91
+ config.setdefault("hf_weight_suffix_dataset", True) # encoder_{dataset}.safetensors etc.
92
+ config.setdefault("hf_preserve_case", True) # keep DatasetName casing in paths
93
+
94
+ # --- Reproducibility / determinism ---
95
+ config.setdefault("deterministic", True) # set cudnn deterministic + disable benchmark
96
+ config.setdefault("strict_determinism", False) # torch.use_deterministic_algorithms(True)
97
+ config.setdefault("deterministic_cublas", False) # set CUBLAS_WORKSPACE_CONFIG
98
+ config.setdefault("seed_per_dataset", False) # re-seed using dataset name in sweep
99
+
100
+
101
+ # ---------------------------------------------------------------------
102
+ # Fast settings / safety
103
+ # ---------------------------------------------------------------------
104
+ torch.autograd.set_detect_anomaly(False)
105
+
106
+ # Determinism knobs (must be set before layers allocate kernels)
107
+ if bool(config.get("deterministic", True)):
108
+ torch.backends.cudnn.benchmark = False
109
+ torch.backends.cudnn.deterministic = True
110
+ else:
111
+ torch.backends.cudnn.benchmark = True
112
+
113
+ # TF32 off → numerically stable & repeatable on Ampere+
114
+ torch.backends.cudnn.allow_tf32 = False
115
+ torch.backends.cuda.matmul.allow_tf32 = False
116
+
117
+ # cuBLAS deterministic workspace (opt-in; can slow some kernels)
118
+ if bool(config.get("deterministic_cublas", False)):
119
+ os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
120
+
121
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
122
+ print(f"Using device: {device}")
123
+
124
+
125
+
126
+ print("\n" + "="*60)
127
+ print("PENTACHORON CONSTELLATION CONFIGURATION")
128
+ print("="*60)
129
+ for k, v in config.items():
130
+ print(f"{k:24}: {v}")
131
+
132
+ # ---------------------------------------------------------------------
133
+ # Reproducibility
134
+ # ---------------------------------------------------------------------
135
+ # ---------------------------------------------------------------------
136
+ # Reproducibility
137
+ # ---------------------------------------------------------------------
138
+ def seed_everything(seed: int = 42,
139
+ deterministic: bool | None = None,
140
+ strict: bool | None = None):
141
+ """Seed Python, NumPy, Torch (CPU+CUDA), and set hash seed/env flags."""
142
+ if deterministic is None:
143
+ deterministic = bool(config.get("deterministic", True))
144
+ if strict is None:
145
+ strict = bool(config.get("strict_determinism", False))
146
+
147
+ # OS / interpreter
148
+ os.environ["PYTHONHASHSEED"] = str(seed)
149
+ try:
150
+ import torch
151
+ torch.use_deterministic_algorithms(strict) # raises on nondet ops if True
152
+ except Exception:
153
+ pass
154
+
155
+ # RNGs
156
+ random.seed(seed)
157
+ np.random.seed(seed)
158
+ torch.manual_seed(seed)
159
+ if torch.cuda.is_available():
160
+ torch.cuda.manual_seed(seed)
161
+ torch.cuda.manual_seed_all(seed)
162
+
163
+ def make_torch_generator(seed: int) -> torch.Generator:
164
+ g = torch.Generator()
165
+ g.manual_seed(seed)
166
+ return g
167
+
168
+ def seed_worker(worker_id: int):
169
+ """Seed DataLoader worker; uses PyTorch's initial_seed to derive unique stream."""
170
+ worker_seed = torch.initial_seed() % 2**32
171
+ np.random.seed(worker_seed)
172
+ random.seed(worker_seed)
173
+
174
+ # Initial global seed
175
+ seed_everything(int(config.get("seed", 42)))
176
+
177
+ # ---------------------------------------------------------------------
178
+ # Setup & deps
179
+ # ---------------------------------------------------------------------
180
+ def _ensure(pkg, pip_name=None):
181
+ pip_name = pip_name or pkg
182
+ try:
183
+ __import__(pkg)
184
+ except Exception:
185
+ print(f"[setup] Installing {pip_name} ...")
186
+ os.system(f"{sys.executable} -m pip install -q {pip_name}")
187
+
188
+ _ensure("safetensors")
189
+ _ensure("huggingface_hub")
190
+ _ensure("pandas")
191
+ _ensure("psutil")
192
+ _ensure("medmnist")
193
+
194
+ from safetensors.torch import save_file as save_safetensors
195
+ from huggingface_hub import HfApi, create_repo, whoami, login
196
+ import pandas as pd
197
+ import psutil
198
+ from torch.utils.tensorboard import SummaryWriter
199
+
200
+ # ---------------------------------------------------------------------
201
+ # Small utils
202
+ # ---------------------------------------------------------------------
203
+ def _timestamp() -> str:
204
+ return datetime.now().strftime("%Y%m%d-%H%M%S")
205
+
206
+ def _param_count(m: nn.Module) -> int:
207
+ return sum(p.numel() for p in m.parameters())
208
+
209
+ def _resolve_repo_id(cfg: Dict) -> str:
210
+ rid = os.getenv("PENTACHORA_HF_REPO") or cfg.get("hf_repo_id")
211
+ if not rid:
212
+ raise RuntimeError("Set config['hf_repo_id'] or export PENTACHORA_HF_REPO.")
213
+ return rid
214
+
215
+ def _hf_login_if_needed():
216
+ try:
217
+ _ = whoami()
218
+ except Exception:
219
+ token = os.getenv("HF_TOKEN")
220
+ if token:
221
+ login(token=token, add_to_git_credential=True)
222
+ else:
223
+ print("[huggingface] No login found and HF_TOKEN not set. Push may fail; run `huggingface-cli login`.")
224
+
225
+ def _ensure_repo(repo_id: str) -> HfApi:
226
+ api = HfApi()
227
+ create_repo(repo_id=repo_id, private=False, exist_ok=True, repo_type="model")
228
+ return api
229
+
230
+ def _zip_dir(src: Path, dst_zip: Path):
231
+ with zipfile.ZipFile(dst_zip, "w", zipfile.ZIP_DEFLATED) as z:
232
+ for p in src.rglob("*"):
233
+ z.write(p, arcname=p.relative_to(src))
234
+
235
+ def _dataset_slug(name_or_names) -> str:
236
+ if isinstance(name_or_names, (list, tuple)):
237
+ return "+".join(n.strip().lower() for n in name_or_names)
238
+ return str(name_or_names).strip().lower()
239
+
240
+ # ---------------------------------------------------------------------
241
+ # Dataset loader (TorchVision + MedMNIST), config-aware
242
+ # ---------------------------------------------------------------------
243
+ try:
244
+ import medmnist
245
+ from medmnist import INFO as MED_INFO
246
+ except Exception:
247
+ medmnist = None
248
+ MED_INFO = None
249
+
250
+ _TORCHVISION_KEYS = {
251
+ "mnist": "MNIST",
252
+ "fashionmnist": "FashionMNIST",
253
+ "kmnist": "KMNIST",
254
+ "emnist": "EMNIST", # balanced
255
+ "qmnist": "QMNIST",
256
+ "usps": "USPS",
257
+ }
258
+ _MEDMNIST_MAP = {
259
+ "bloodmnist": "bloodmnist", "pathmnist": "pathmnist", "octmnist": "octmnist",
260
+ "pneumoniamnist": "pneumoniamnist", "dermamnist": "dermamnist", "retinamnist": "retinamnist",
261
+ "breastmnist": "breastmnist", "organamnist": "organamnist", "organcmnist": "organcmnist",
262
+ "organsmnist": "organsmnist", "tissuemnist": "tissuemnist",
263
+ }
264
+
265
+ _DATASET_STATS_1CH = {
266
+ "MNIST": ([0.1307], [0.3081]),
267
+ "FashionMNIST": ([0.2860], [0.3530]),
268
+ "KMNIST": ([0.1918], [0.3483]),
269
+ "EMNIST": ([0.1307], [0.3081]),
270
+ "QMNIST": ([0.1307], [0.3081]),
271
+ "USPS": ([0.5000], [0.5000]),
272
+ }
273
+ _MEAN1, _STD1 = [0.5], [0.5]
274
+ _MEAN3, _STD3 = [0.5, 0.5, 0.5], [0.5, 0.5, 0.5]
275
+
276
+ def _norm_stats(name: str, channels: int) -> Tuple[List[float], List[float]]:
277
+ if channels == 1:
278
+ return _DATASET_STATS_1CH.get(name, (_MEAN1, _STD1))
279
+ return _MEAN3, _STD3
280
+
281
+ def _to_channels(target_c: int):
282
+ def _fn(t: torch.Tensor) -> torch.Tensor:
283
+ c = t.shape[0]
284
+ if c == target_c:
285
+ return t
286
+ if target_c == 1:
287
+ if c == 3:
288
+ r, g, b = t[0], t[1], t[2]
289
+ gray = 0.2989*r + 0.5870*g + 0.1140*b
290
+ return gray.unsqueeze(0)
291
+ return t[:1]
292
+ if target_c == 3:
293
+ if c == 1:
294
+ return t.repeat(3, 1, 1)
295
+ return t[:3]
296
+ return t[:target_c]
297
+ return transforms.Lambda(_fn)
298
+
299
+ def _augmentations_for(name: str, size: int, channels: int) -> List[transforms.Transform]:
300
+ aug = []
301
+ if not bool(config.get("augment", False)):
302
+ return aug
303
+ if name.upper() in {"MNIST","KMNIST","EMNIST","QMNIST","USPS"}:
304
+ aug += [transforms.RandomAffine(degrees=8, translate=(0.05, 0.05), scale=(0.95, 1.05))]
305
+ if size >= 32:
306
+ pad = max(1, int(0.03 * size))
307
+ aug += [transforms.RandomCrop(size, padding=pad)]
308
+ return aug
309
+ if size >= 32:
310
+ pad = max(1, int(0.03 * size))
311
+ aug += [transforms.RandomCrop(size, padding=pad)]
312
+ aug += [transforms.RandomAffine(degrees=10, translate=(0.05, 0.05), scale=(0.95, 1.05))]
313
+ if channels == 3 and name.lower().endswith("mnist"):
314
+ aug += [transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.05, hue=0.02)]
315
+ return aug
316
+
317
+ def _build_transforms(dataset_name: str, split: str, native_c: int, target_c: int|str, size: int, normalize: bool, per_dataset_norm: bool) -> transforms.Compose:
318
+ t: List[transforms.Transform] = []
319
+ if size != 28:
320
+ t.append(transforms.Resize((size, size)))
321
+ t.append(transforms.ToTensor())
322
+ out_c = native_c
323
+ if target_c != "auto":
324
+ t.append(_to_channels(int(target_c)))
325
+ out_c = int(target_c)
326
+ if split == "train":
327
+ t = _augmentations_for(dataset_name, size, out_c) + t
328
+ if normalize:
329
+ if per_dataset_norm:
330
+ mean, std = _norm_stats(dataset_name, out_c)
331
+ else:
332
+ mean, std = (_MEAN1, _STD1) if out_c == 1 else (_MEAN3, _STD3)
333
+ t.append(transforms.Normalize(mean=mean, std=std))
334
+ t.append(transforms.Lambda(lambda x: x.view(-1)))
335
+ return transforms.Compose(t)
336
+
337
+ def collate_as_int(batch):
338
+ xs, ys = zip(*batch)
339
+ xs = torch.stack(xs, dim=0)
340
+ _ys = []
341
+ for y in ys:
342
+ if isinstance(y, (int, np.integer)):
343
+ _ys.append(int(y))
344
+ elif torch.is_tensor(y):
345
+ if y.ndim == 0: _ys.append(int(y.item()))
346
+ elif y.ndim == 1 and y.numel()==1: _ys.append(int(y.item()))
347
+ else: _ys.append(int(y.argmax().item()))
348
+ else:
349
+ arr = np.asarray(y)
350
+ if arr.ndim == 0 or (arr.ndim==1 and arr.size==1):
351
+ _ys.append(int(arr.item()))
352
+ else:
353
+ _ys.append(int(arr.argmax()))
354
+ ys_tensor = torch.tensor(_ys, dtype=torch.long)
355
+ return xs, ys_tensor
356
+
357
+ def _get_med_info(flag: str) -> dict:
358
+ if MED_INFO is None:
359
+ raise ImportError("medmnist is not installed. `pip install medmnist`")
360
+ if flag not in MED_INFO:
361
+ raise KeyError(f"Unknown MedMNIST flag: {flag}")
362
+ return MED_INFO[flag]
363
+
364
+ def _med_class_names(info: dict) -> List[str]:
365
+ lab = info["label"]
366
+ return [lab[str(i)] for i in range(len(lab))]
367
+
368
+ def load_single_dataset(name: str, split: str,
369
+ cfg: Optional[Dict]=None,
370
+ resolved_target_channels: Optional[int|str]=None
371
+ ) -> Tuple[torch.utils.data.Dataset, int, List[str], int, int]:
372
+ """
373
+ Return: dataset, num_classes, class_names, input_dim (C*H*W), output_channels
374
+ """
375
+ cfg = cfg or config
376
+ name_key = name.strip()
377
+ name_lower = name_key.lower()
378
+
379
+ size = int(cfg.get("img_size", 28))
380
+ want_c = cfg.get("img_channels", "auto") if resolved_target_channels is None else resolved_target_channels
381
+ normalize = bool(cfg.get("normalize", True))
382
+ per_dataset_norm = bool(cfg.get("per_dataset_norm", True))
383
+
384
+ # TorchVision
385
+ if name_lower in _TORCHVISION_KEYS:
386
+ canonical = _TORCHVISION_KEYS[name_lower]
387
+ native_c = 1
388
+ transform = _build_transforms(canonical, split, native_c, want_c, size, normalize, per_dataset_norm)
389
+
390
+ if canonical == "MNIST":
391
+ ds = datasets.MNIST("./data", train=(split=="train"), transform=transform, download=True)
392
+ ncls = 10; cls_names = [f"digit-{i}" for i in range(10)]
393
+ elif canonical == "FashionMNIST":
394
+ base = ['T-shirt/top','Trouser','Pullover','Dress','Coat','Sandal','Shirt','Sneaker','Bag','Ankle boot']
395
+ ds = datasets.FashionMNIST("./data", train=(split=="train"), transform=transform, download=True)
396
+ ncls = 10; cls_names = [f"fashion-{n}" for n in base]
397
+ elif canonical == "KMNIST":
398
+ ds = datasets.KMNIST("./data", train=(split=="train"), transform=transform, download=True)
399
+ ncls = 10; cls_names = [f"kmnist-{c}" for c in ['お','き','す','つ','な','は','ま','や','れ','を']]
400
+ elif canonical == "EMNIST":
401
+ ds = datasets.EMNIST("./data", split='balanced', train=(split=="train"), transform=transform, download=True)
402
+ letters = ['0','1','2','3','4','5','6','7','8','9',
403
+ 'A','B','C','D','E','F','G','H','I','J','K','L','M','N','O','P','Q','R','S','T','U','V','W','X','Y','Z',
404
+ 'a','b','d','e','f','g','h','n','q','r','t']
405
+ ncls = 47; cls_names = [f"emnist-{c}" for c in letters]
406
+ elif canonical == "QMNIST":
407
+ ds = datasets.QMNIST("./data", what=('train' if split=="train" else 'test'), transform=transform, download=True)
408
+ ncls = 10; cls_names = [f"qmnist-{i}" for i in range(10)]
409
+ elif canonical == "USPS":
410
+ ds = datasets.USPS("./data", train=(split=="train"), transform=transform, download=True)
411
+ ncls = 10; cls_names = [f"usps-{i}" for i in range(10)]
412
+ else:
413
+ raise ValueError(f"Unhandled TorchVision dataset: {canonical}")
414
+
415
+ out_c = native_c if want_c == "auto" else int(want_c)
416
+ input_dim = out_c * size * size
417
+ return ds, ncls, cls_names, input_dim, out_c
418
+
419
+ # MedMNIST
420
+ if name_lower in _MEDMNIST_MAP:
421
+ if medmnist is None:
422
+ raise ImportError("medmnist not available. `pip install medmnist`")
423
+ flag = _MEDMNIST_MAP[name_lower]
424
+ info = _get_med_info(flag)
425
+ DataClass = getattr(medmnist, info["python_class"])
426
+ native_c = int(info.get("n_channels", 1))
427
+ out_c = native_c if want_c == "auto" else int(want_c)
428
+
429
+ transform = transforms.Compose([
430
+ transforms.ToTensor(),
431
+ _to_channels(out_c) if want_c != "auto" else transforms.Lambda(lambda t: t),
432
+ *(_augmentations_for(name_key, size, out_c) if (split=="train") else []),
433
+ transforms.Resize((size, size)) if size != 28 else transforms.Lambda(lambda t: t),
434
+ transforms.Normalize(*(_norm_stats(name_key, out_c) if (normalize and per_dataset_norm) else ((_MEAN1,_STD1) if out_c==1 else (_MEAN3,_STD3)))) if normalize else transforms.Lambda(lambda t: t),
435
+ transforms.Lambda(lambda t: t.view(-1)),
436
+ ])
437
+ ds = DataClass(split=('train' if split=="train" else 'test'), transform=transform, download=True, size=size)
438
+ ncls = len(info["label"]); cls_names = _med_class_names(info)
439
+ input_dim = out_c * size * size
440
+ return ds, ncls, cls_names, input_dim, out_c
441
+
442
+ raise ValueError(f"Unknown dataset name: {name}")
443
+
444
+ def get_dataset_single(name: str, batch_size: int, num_workers: int = 2):
445
+ """
446
+ Load a single dataset honoring config overrides.
447
+ Returns: train_loader, test_loader, num_classes, class_names, input_dim, channels
448
+ """
449
+ tr, ntr, names_tr, in_tr, out_c = load_single_dataset(name, "train", config)
450
+ te, nte, names_te, in_te, out_c2 = load_single_dataset(name, "test", config)
451
+ assert ntr == nte and in_tr == in_te and out_c == out_c2
452
+ g = make_torch_generator(int(config.get("seed", 42)))
453
+ train_loader = DataLoader(
454
+ tr, batch_size=batch_size, shuffle=True, num_workers=num_workers,
455
+ pin_memory=torch.cuda.is_available(), collate_fn=collate_as_int,
456
+ worker_init_fn=seed_worker, generator=g, persistent_workers=False
457
+ )
458
+ test_loader = DataLoader(
459
+ te, batch_size=batch_size, shuffle=False, num_workers=num_workers,
460
+ pin_memory=torch.cuda.is_available(), collate_fn=collate_as_int,
461
+ worker_init_fn=seed_worker, generator=g, persistent_workers=False
462
+ )
463
+
464
+ return train_loader, test_loader, ntr, names_tr, in_tr, out_c
465
+
466
+ # Dataset catalogs
467
+ TORCHVISION_DATASETS = ["MNIST", "FashionMNIST", "KMNIST", "EMNIST", "QMNIST", "USPS"]
468
+ MEDMNIST_DATASETS = ["BloodMNIST","PathMNIST","OCTMNIST","PneumoniaMNIST","DermaMNIST",
469
+ "RetinaMNIST","BreastMNIST","OrganAMNIST","OrganCMNIST","OrganSMNIST","TissueMNIST"]
470
+
471
+ # ---------------------------------------------------------------------
472
+ # Models
473
+ # ---------------------------------------------------------------------
474
+ class PentaFreqExtractor(nn.Module):
475
+ """
476
+ Multi-channel spectral extractor:
477
+ - Input: [B, C*H*W], unflatten -> [B, C, H, W]
478
+ - 5 frequency bands -> encode to base_dim each
479
+ """
480
+ def __init__(self, input_dim: int = 784, input_ch: int = 1, base_dim: int = 64, channels: int = 12):
481
+ super().__init__()
482
+ self.input_dim = input_dim
483
+ self.input_ch = int(input_ch)
484
+ side_f = (input_dim / max(1, self.input_ch)) ** 0.5
485
+ side = int(side_f)
486
+ assert side * side * self.input_ch == input_dim, f"input_dim ({input_dim}) != C*H*W with H=W; C={self.input_ch}, side≈{side_f:.3f}"
487
+ self.unflatten = nn.Unflatten(1, (self.input_ch, side, side))
488
+ self.base_dim = base_dim
489
+
490
+ # Vertex 0 (ultra-high)
491
+ self.v0_ultrahigh = nn.Sequential(
492
+ nn.Conv2d(self.input_ch, channels, 3, padding=1),
493
+ nn.BatchNorm2d(channels), nn.ReLU(),
494
+ nn.Conv2d(channels, channels, 3, padding=1, groups=channels),
495
+ nn.BatchNorm2d(channels), nn.ReLU(),
496
+ nn.AdaptiveAvgPool2d(7), nn.Flatten()
497
+ ); self.v0_encode = nn.Linear(channels * 49, base_dim)
498
+
499
+ # Vertex 1 (high)
500
+ self.v1_high = nn.Sequential(
501
+ nn.Conv2d(self.input_ch, channels, 3, padding=1),
502
+ nn.BatchNorm2d(channels), nn.Tanh(),
503
+ nn.MaxPool2d(2),
504
+ nn.Conv2d(channels, channels, 3, padding=1),
505
+ nn.BatchNorm2d(channels),
506
+ nn.Tanh(),
507
+ nn.AdaptiveAvgPool2d(7),
508
+ nn.Flatten()
509
+ ); self.v1_encode = nn.Linear(channels * 49, base_dim)
510
+
511
+ # Vertex 2 (mid)
512
+ self.v2_mid = nn.Sequential(
513
+ nn.Conv2d(self.input_ch, channels, 5, padding=2, stride=2),
514
+ nn.BatchNorm2d(channels), nn.GELU(),
515
+ nn.Conv2d(channels, channels, 3, padding=1),
516
+ nn.BatchNorm2d(channels),
517
+ nn.GELU(),
518
+ nn.AdaptiveAvgPool2d(7),
519
+ nn.Flatten()
520
+ ); self.v2_encode = nn.Linear(channels * 49, base_dim)
521
+
522
+ # Vertex 3 (low-mid)
523
+ self.v3_lowmid = nn.Sequential(
524
+ nn.AvgPool2d(2),
525
+ nn.Conv2d(self.input_ch, channels, 7, padding=3),
526
+ nn.BatchNorm2d(channels), nn.SiLU(),
527
+ nn.AdaptiveAvgPool2d(7), nn.Flatten()
528
+ ); self.v3_encode = nn.Linear(channels * 49, base_dim)
529
+
530
+ # Vertex 4 (low)
531
+ self.v4_low = nn.Sequential(
532
+ nn.AvgPool2d(4),
533
+ nn.Conv2d(self.input_ch, channels, 7, padding=3),
534
+ nn.BatchNorm2d(channels),
535
+ nn.Sigmoid(),
536
+ nn.AdaptiveAvgPool2d(7),
537
+ nn.Flatten()
538
+ ); self.v4_encode = nn.Linear(channels * 49, base_dim)
539
+
540
+ self.register_buffer("adjacency_matrix", torch.ones(5, 5) - torch.eye(5))
541
+ self._init_edge_kernels(channels)
542
+
543
+ @torch.no_grad()
544
+ def _init_edge_kernels(self, channels: int):
545
+ if channels < 5: return
546
+ conv0 = self.v0_ultrahigh[0]
547
+ if not isinstance(conv0, nn.Conv2d): return
548
+ if conv0.weight.shape[1] >= 1:
549
+ k = conv0.weight
550
+ k[0,0] = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=k.dtype)/2
551
+ k[1,0] = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]], dtype=k.dtype)/2
552
+ k[2,0] = torch.tensor([[0,-1,0],[-1,4,-1],[0,-1,0]], dtype=k.dtype)/3
553
+ k[3,0] = torch.tensor([[1,0,0],[0,-1,0],[0,0,0]], dtype=k.dtype)/4
554
+ k[4,0] = torch.tensor([[-1,0,1],[-1,0,1],[-1,0,1]], dtype=k.dtype)/5
555
+
556
+ def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
557
+ img = self.unflatten(x)
558
+ v0 = self.v0_encode(self.v0_ultrahigh(img))
559
+ v1 = self.v1_encode(self.v1_high(img))
560
+ v2 = self.v2_encode(self.v2_mid(img))
561
+ v3 = self.v3_encode(self.v3_lowmid(img))
562
+ v4 = self.v4_encode(self.v4_low(img))
563
+ vertices = torch.stack([v0, v1, v2, v3, v4], dim=1) # [B,5,D]
564
+ return vertices, self.adjacency_matrix
565
+
566
+ class PentachoronCrossAttention(nn.Module):
567
+ def __init__(self, dim: int, num_heads: int = 14, dropout: float = 0.0):
568
+ super().__init__()
569
+ self.attn = nn.MultiheadAttention(dim, num_heads=num_heads, dropout=dropout, batch_first=True)
570
+ def _row_to_attn_mask(self, row: torch.Tensor) -> torch.Tensor:
571
+ mask = torch.zeros(1, row.numel(), device=row.device, dtype=torch.float32)
572
+ mask[(row == 0).unsqueeze(0)] = float("-inf")
573
+ return mask
574
+ def forward(self, vertices: torch.Tensor, adjacency: torch.Tensor) -> torch.Tensor:
575
+ B, V, D = vertices.shape
576
+ outs = []
577
+ for i in range(V):
578
+ q = vertices[:, i:i+1, :]
579
+ k = v = vertices
580
+ mask = self._row_to_attn_mask(adjacency[i].to(vertices.device))
581
+ out, _ = self.attn(q, k, v, attn_mask=mask, need_weights=False)
582
+ outs.append(out)
583
+ return torch.cat(outs, dim=1)
584
+
585
+ class PentachoronOpinionFusion(nn.Module):
586
+ def __init__(self, base_dim: int = 64, proj_dim: Optional[int] = None, num_heads: int = 14, p_dropout: float = 0.0):
587
+ super().__init__()
588
+ self.cross = PentachoronCrossAttention(dim=base_dim, num_heads=num_heads)
589
+ self.fusion = nn.Sequential(
590
+ nn.Linear(base_dim * 5, base_dim * 4),
591
+ nn.BatchNorm1d(base_dim * 4),
592
+ nn.ReLU(),
593
+ nn.Dropout(p_dropout),
594
+ nn.Linear(base_dim * 4, base_dim * 3),
595
+ nn.BatchNorm1d(base_dim * 3),
596
+ nn.ReLU(),
597
+ nn.Dropout(p_dropout),
598
+ nn.Linear(base_dim * 3, base_dim * 2),
599
+ nn.BatchNorm1d(base_dim * 2),
600
+ nn.ReLU(),
601
+ nn.Linear(base_dim * 2, base_dim),
602
+ )
603
+ self.projection = None if proj_dim is None else nn.Linear(base_dim, proj_dim, bias=False)
604
+ self._lambda_raw = nn.Parameter(torch.tensor(0.5))
605
+
606
+ @staticmethod
607
+ def _softmax_geometry(vertices: torch.Tensor, adjacency: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
608
+ v_norm = F.normalize(vertices, dim=2, eps=1e-8)
609
+ sims = torch.bmm(v_norm, v_norm.transpose(1, 2))
610
+ edge_strengths = sims * adjacency.to(vertices.dtype).unsqueeze(0)
611
+ weights = F.softmax(edge_strengths.sum(dim=2), dim=1) # [B,5]
612
+ weighted = vertices * weights.unsqueeze(2)
613
+ return weighted, weights
614
+
615
+ def forward(self, vertices: torch.Tensor, adjacency: torch.Tensor, return_diag: bool = False):
616
+ soft_out, weights = self._softmax_geometry(vertices, adjacency)
617
+ attn_out = self.cross(vertices, adjacency)
618
+ lam = torch.sigmoid(self._lambda_raw)
619
+ combined = lam * soft_out + (1.0 - lam) * attn_out
620
+ fused = self.fusion(combined.flatten(1))
621
+ if self.projection is not None:
622
+ fused = self.projection(fused)
623
+ z = F.normalize(fused, dim=1)
624
+ if not return_diag:
625
+ return z, None
626
+ return z, {"lambda": lam.detach(), "softmax_weights": weights.detach()}
627
+
628
+ class PentaFreqEncoderV2(nn.Module):
629
+ def __init__(self, input_dim: int = 784, input_ch: int = 1, base_dim: int = 64, proj_dim: Optional[int] = None, num_heads: int = 14, channels: int = 12):
630
+ super().__init__()
631
+ self.extractor = PentaFreqExtractor(input_dim=input_dim, input_ch=input_ch, base_dim=base_dim, channels=channels)
632
+ self.opinion = PentachoronOpinionFusion(base_dim=base_dim, proj_dim=proj_dim, num_heads=num_heads)
633
+ @torch.no_grad()
634
+ def get_frequency_contributions(self, x: torch.Tensor) -> torch.Tensor:
635
+ verts, adj = self.extractor(x)
636
+ _, w = self.opinion._softmax_geometry(verts, adj)
637
+ return w
638
+ def forward(self, x: torch.Tensor, return_diag: bool = False):
639
+ verts, adj = self.extractor(x)
640
+ z, diag = self.opinion(verts, adj, return_diag)
641
+ return (z, diag) if return_diag else z
642
+
643
+
644
+ from geovocab2.shapes.factory.simplex_factory import SimplexFactory
645
+
646
+ class BatchedPentachoronConstellation(nn.Module):
647
+ def __init__(self, num_classes: int, dim: int, num_pairs: int = 5, device: Optional[torch.device] = None, lambda_sep: float = 0.5, num_heads=14):
648
+ super().__init__()
649
+ self.num_heads = num_heads
650
+ self.num_classes = num_classes
651
+ self.dim = dim
652
+ self.num_pairs = num_pairs
653
+ self.device = device if device is not None else torch.device("cpu")
654
+ self.lambda_separation = lambda_sep
655
+
656
+ self.dispatchers = nn.Parameter(self._init_batched_pentachora())
657
+ self.specialists = nn.Parameter(self._init_batched_pentachora())
658
+
659
+ self.dispatcher_weights = nn.Parameter(torch.randn(num_pairs, 5) * 0.1)
660
+ self.specialist_weights = nn.Parameter(torch.randn(num_pairs, 5) * 0.1)
661
+ self.temps = nn.Parameter(lambda_sep * torch.ones(num_pairs))
662
+
663
+ self.register_buffer("vertex_map", self._create_vertex_mapping())
664
+
665
+ self.group_heads = nn.ModuleList([
666
+ nn.Linear(dim, int((self.vertex_map == i).sum().item())) if int((self.vertex_map == i).sum().item()) > 0 else None
667
+ for i in range(5)
668
+ ])
669
+
670
+ self.cross_attention = nn.MultiheadAttention(embed_dim=dim, num_heads=self.num_heads, dropout=0.0, batch_first=True)
671
+ self.aggregation_weights = nn.Parameter(torch.ones(num_pairs) / num_pairs)
672
+
673
+ self.fusion = nn.Sequential(
674
+ nn.Linear(num_classes * num_pairs, num_classes * 2),
675
+ nn.BatchNorm1d(num_classes * 2),
676
+ nn.ReLU(),
677
+ nn.Dropout(0.0),
678
+ nn.Linear(num_classes * 2, num_classes)
679
+ )
680
+
681
+ self.coherence_head = nn.Sequential(nn.Linear(dim, dim // 2), nn.GELU(), nn.Linear(dim // 2, 1))
682
+
683
+ self.blend_weight = nn.Parameter(torch.tensor(0.5))
684
+
685
+ #def _init_batched_pentachora(self) -> torch.Tensor:
686
+ # sqrt15, sqrt10, sqrt5 = np.sqrt(15), np.sqrt(10), np.sqrt(5)
687
+ # base_simplex = torch.tensor([
688
+ # [ 1.0, 0.0, 0.0, 0.0],
689
+ # [-0.25, sqrt15/4, 0.0, 0.0],
690
+ # [-0.25,-sqrt15/12, sqrt10/3, 0.0],
691
+ # [-0.25,-sqrt15/12,-sqrt10/6, sqrt5/2],
692
+ # [-0.25,-sqrt15/12,-sqrt10/6,-sqrt5/2]
693
+ # ], device=self.device, dtype=torch.float32)
694
+ # base_simplex = F.normalize(base_simplex, dim=1)
695
+ # pentachora = torch.zeros(self.num_pairs, 5, self.dim, device=self.device, dtype=torch.float32)
696
+ # for i in range(self.num_pairs):
697
+ # pentachora[i, :, :4] = base_simplex * (1 + 0.1 * i)
698
+ # if self.dim > 4:
699
+ # pentachora[i, :, 4:] = torch.randn(5, self.dim - 4, device=self.device) * (random.random() * 0.25)
700
+ # return pentachora * 2.0
701
+
702
+ # In your constellation file, add at the top:
703
+
704
+
705
+ from geovocab2.shapes.factory.simplex_factory import SimplexFactory
706
+
707
+ def _init_batched_pentachora(self) -> torch.Tensor:
708
+ """Initialize pentachora using SimplexFactory for stable geometry."""
709
+ factory = SimplexFactory(
710
+ k=4, # 4-simplex = pentachoron
711
+ embed_dim=self.dim,
712
+ method="regular", # equal edges, maximal volume
713
+ scale=2.0
714
+ )
715
+
716
+ pentachora = torch.zeros(self.num_pairs, 5, self.dim, device=self.device, dtype=torch.float32)
717
+
718
+ for i in range(self.num_pairs):
719
+ base = factory.build(
720
+ backend="torch",
721
+ device=str(self.device),
722
+ seed=42 + i,
723
+ validate=True
724
+ )
725
+ #pentachora[i] = base * (1.0 + 0.1 * i)
726
+
727
+ return pentachora
728
+
729
+ def _create_vertex_mapping(self) -> torch.Tensor:
730
+ mapping = torch.zeros(self.num_classes, dtype=torch.long)
731
+ for i in range(self.num_classes):
732
+ mapping[i] = i % 5
733
+ return mapping
734
+
735
+ def forward(self, x: torch.Tensor):
736
+ B = x.size(0)
737
+ coherence_gate = torch.sigmoid(self.coherence_head(x)) # [B,1]
738
+
739
+ x_exp = x.unsqueeze(1).unsqueeze(2) # [B,1,1,D]
740
+ disp_exp = self.dispatchers.unsqueeze(0) # [1,P,5,D]
741
+ spec_exp = self.specialists.unsqueeze(0) # [1,P,5,D]
742
+ disp_d = torch.norm(x_exp - disp_exp, dim=3) # [B,P,5]
743
+ spec_d = torch.norm(x_exp - spec_exp, dim=3) # [B,P,5]
744
+
745
+ dw = F.softmax(self.dispatcher_weights, dim=1).unsqueeze(0)
746
+ sw = F.softmax(self.specialist_weights, dim=1).unsqueeze(0)
747
+ temps = torch.clamp(self.temps, 0.1, 2.0).view(1, -1, 1)
748
+
749
+ disp_logits = -(disp_d * dw) / temps
750
+ spec_logits = -(spec_d * sw) / temps
751
+
752
+ c = coherence_gate.unsqueeze(-1)
753
+ disp_probs = F.softmax(disp_logits * c, dim=2)
754
+ spec_probs = F.softmax(spec_logits * c, dim=2)
755
+ probs = 0.5 * disp_probs + 0.5 * spec_probs
756
+
757
+ scores_by_pair = []
758
+ for p in range(self.num_pairs):
759
+ pair_scores = torch.zeros(B, self.num_classes, device=x.device)
760
+ for v_idx in range(5):
761
+ idxs = (self.vertex_map == v_idx).nonzero(as_tuple=True)[0]
762
+ if len(idxs) == 0: continue
763
+ v_prob = probs[:, p, v_idx:v_idx+1]
764
+ if self.group_heads[v_idx] is not None:
765
+ g_logits = self.group_heads[v_idx](x) # [B, |idxs|]
766
+ gated = g_logits * v_prob
767
+ for i, cls in enumerate(idxs.tolist()):
768
+ if i < gated.size(1):
769
+ pair_scores[:, cls] = gated[:, i]
770
+ scores_by_pair.append(pair_scores)
771
+
772
+ scores_tensor = torch.stack(scores_by_pair, dim=1) # [B,P,C]
773
+
774
+ centers = self.dispatchers.mean(dim=1).unsqueeze(0).expand(B, -1, -1)
775
+ _attn, _ = self.cross_attention(centers, centers, centers)
776
+
777
+ agg = F.softmax(self.aggregation_weights, dim=0).view(1, -1, 1)
778
+ weighted = (scores_tensor * agg).sum(dim=1) # [B,C]
779
+ fused = self.fusion(scores_tensor.flatten(1)) # [B,C]
780
+ alpha = torch.sigmoid(self.blend_weight)
781
+ final = alpha * weighted + (1 - alpha) * fused
782
+ return final, {"disp_d": disp_d, "spec_d": spec_d, "probs": probs}
783
+
784
+ def _batched_cayley_menger(self, pentachora: torch.Tensor) -> torch.Tensor:
785
+ """
786
+ Stable CM proxy: det(M) via eigvals of (M + eps*I).
787
+ Returns a positive scalar per cube; larger => more 'volumetric' (less degenerate).
788
+ """
789
+ P = pentachora.shape[0]
790
+ d2 = torch.cdist(pentachora, pentachora).pow(2) # [P,5,5]
791
+ M = torch.zeros(P, 6, 6, device=self.device, dtype=pentachora.dtype)
792
+ M[:, 0, 1:] = 1.0
793
+ M[:, 1:, 0] = 1.0
794
+ M[:, 1:, 1:] = d2
795
+
796
+ eps = 1e-6
797
+ I = torch.eye(6, device=self.device, dtype=pentachora.dtype).unsqueeze(0)
798
+ M_eps = M + eps * I # make SPD-ish
799
+ # eigvalsh → real, sorted
800
+ evals = torch.linalg.eigvalsh(M_eps) # [P,6]
801
+ evals = evals.clamp_min(1e-12) # avoid log(<=0)
802
+ logdet = evals.log().sum(dim=1) # log|det|
803
+ det = torch.exp(logdet) # |det|
804
+ # keep it finite
805
+ det = torch.nan_to_num(det, nan=0.0, posinf=1e6, neginf=0.0)
806
+ return det
807
+
808
+
809
+ def _batched_edge_variance(self, pentachora: torch.Tensor) -> torch.Tensor:
810
+ d = torch.cdist(pentachora, pentachora)
811
+ mask = torch.triu(torch.ones(5, 5, device=pentachora.device), diagonal=1).bool()
812
+ edges = torch.stack([d[p][mask] for p in range(self.num_pairs)]) # [P,10]
813
+ return edges.var(dim=1).sum() + torch.relu(0.5 - edges.min(dim=1)[0]).sum()
814
+
815
+ def regularization_loss(self, vertex_weights=None) -> torch.Tensor:
816
+ disp_cm = self._batched_cayley_menger(self.dispatchers)
817
+ spec_cm = self._batched_cayley_menger(self.specialists)
818
+ cm_loss = torch.relu(1.0 - torch.abs(disp_cm)).sum() + torch.relu(1.0 - torch.abs(spec_cm)).sum()
819
+ edge_loss = self._batched_edge_variance(self.dispatchers) + self._batched_edge_variance(self.specialists)
820
+ disp_centers = self.dispatchers.mean(dim=1)
821
+ spec_centers = self.specialists.mean(dim=1)
822
+ cos_sims = F.cosine_similarity(disp_centers, spec_centers, dim=1, eps=1e-8)
823
+ ortho = torch.abs(cos_sims).sum() * self.lambda_separation
824
+ separations = torch.norm(disp_centers - spec_centers, dim=1)
825
+ sep = torch.relu(2.0 - separations).sum() * self.lambda_separation
826
+
827
+ dyn = 0.0
828
+ if vertex_weights is not None:
829
+ vw = vertex_weights.to(self.dispatchers.device)
830
+ disp_norms = torch.norm(self.dispatchers, p=2, dim=2)
831
+ spec_norms = torch.norm(self.specialists, p=2, dim=2)
832
+ dyn = 0.1 * ((disp_norms * vw.unsqueeze(0)).mean() + (spec_norms * vw.unsqueeze(0)).mean())
833
+
834
+ return ((cm_loss + edge_loss + ortho + sep) / 4.0) / self.num_pairs + dyn
835
+
836
+ # ---------------------------------------------------------------------
837
+ # Losses
838
+ # ---------------------------------------------------------------------
839
+ def dual_contrastive_loss(latents, targets, constellation, temp: float):
840
+ B = latents.size(0)
841
+ z = F.normalize(latents, dim=1, eps=1e-8)
842
+ disp = F.normalize(constellation.dispatchers, dim=2, eps=1e-8)
843
+ spec = F.normalize(constellation.specialists, dim=2, eps=1e-8)
844
+
845
+ # Compute similarity logits: [B, P, V]
846
+ disp_logits = torch.einsum('bd,pvd->bpv', z, disp) / temp
847
+ spec_logits = torch.einsum('bd,pvd->bpv', z, spec) / temp
848
+
849
+ # Get target vertex indices: [B]
850
+ tvert = constellation.vertex_map[targets]
851
+
852
+ # Expand for gathering: [B, P, 1]
853
+ idx = tvert.view(B, 1, 1).expand(B, disp_logits.size(1), 1)
854
+
855
+ # Gather positive logits: [B, P]
856
+ disp_pos = disp_logits.gather(2, idx).squeeze(2)
857
+ spec_pos = spec_logits.gather(2, idx).squeeze(2)
858
+
859
+ # LogSumExp over vertices: [B, P]
860
+ disp_lse = torch.logsumexp(disp_logits, dim=2)
861
+ spec_lse = torch.logsumexp(spec_logits, dim=2)
862
+
863
+ # InfoNCE loss per pair: [B, P]
864
+ disp_loss_per_pair = disp_lse - disp_pos
865
+ spec_loss_per_pair = spec_lse - spec_pos
866
+
867
+ # Sum over pairs (not mean), then mean over batch
868
+ # This prevents dilution with more pairs
869
+ disp_loss = disp_loss_per_pair.sum(dim=1).mean()
870
+ spec_loss = spec_loss_per_pair.sum(dim=1).mean()
871
+
872
+ return disp_loss + spec_loss
873
+
874
+ class RoseDiagnosticHead(nn.Module):
875
+ def __init__(self, latent_dim: int, hidden_dim: int = 128):
876
+ super().__init__()
877
+ self.net = nn.Sequential(nn.Linear(latent_dim, hidden_dim), nn.GELU(), nn.LayerNorm(hidden_dim), nn.Linear(hidden_dim, 1))
878
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
879
+ return self.net(x)
880
+
881
+ def rose_score_magnitude(x, need, relation, purpose, eps: float = 1e-8):
882
+ x_n = F.normalize(x, dim=-1, eps=eps)
883
+ n_n = F.normalize(need, dim=-1, eps=eps)
884
+ r_n = F.normalize(relation, dim=-1, eps=eps)
885
+ p_n = F.normalize(purpose, dim=-1, eps=eps)
886
+ r7 = ((x_n*n_n).sum(-1) + (x_n*r_n).sum(-1) + (x_n*p_n).sum(-1)) / 3.0
887
+ r8 = x.norm(dim=-1).clamp_min(eps)
888
+ return r7 * r8
889
+
890
+ def rose_contrastive_loss(latents, targets, constellation, temp: float = 0.5):
891
+ B, D = latents.shape
892
+ tvert = constellation.vertex_map[targets]
893
+ need = constellation.specialists[:, tvert, :].mean(dim=0)
894
+ relation = constellation.dispatchers[:, tvert, :].mean(dim=0)
895
+ purpose = constellation.specialists.mean(dim=(0, 1)).unsqueeze(0).expand(B, D)
896
+ rose = rose_score_magnitude(latents, need, relation, purpose)
897
+ weights = (1.0 - torch.tanh(rose)).detach()
898
+ spec = F.normalize(constellation.specialists.mean(dim=0), dim=1, eps=1e-8)
899
+ disp = F.normalize(constellation.dispatchers.mean(dim=0), dim=1, eps=1e-8)
900
+ z = F.normalize(latents, dim=1, eps=1e-8)
901
+ spec_logits = (z @ spec.T) / temp
902
+ disp_logits = (z @ disp.T) / temp
903
+ spec_pos = spec_logits.gather(1, tvert.view(-1,1)).squeeze(1)
904
+ disp_pos = disp_logits.gather(1, tvert.view(-1,1)).squeeze(1)
905
+ spec_lse = torch.logsumexp(spec_logits, dim=1)
906
+ disp_lse = torch.logsumexp(disp_logits, dim=1)
907
+ per_sample = temp * ((spec_lse - spec_pos) + (disp_lse - disp_pos))
908
+ return (per_sample * weights).mean(), rose.detach()
909
+
910
+ # ---------------------------------------------------------------------
911
+ # Regularization helpers
912
+ # ---------------------------------------------------------------------
913
+ def get_class_similarity(constellation_model: BatchedPentachoronConstellation, num_classes: int) -> torch.Tensor:
914
+ W = constellation_model.fusion[-1].weight.data.detach()
915
+ Wn = F.normalize(W, p=2, dim=1)
916
+ return torch.clamp(Wn @ Wn.T, 0.0, 1.0)
917
+
918
+ def vertex_weights_from_confusion(cm: np.ndarray, class_similarity: torch.Tensor, vertex_map: torch.Tensor, device: torch.device) -> torch.Tensor:
919
+ C = cm.shape[0]
920
+ totals = cm.sum(axis=1)
921
+ correct = cm.diagonal()
922
+ acc = np.divide(correct, totals, out=np.zeros_like(correct, dtype=float), where=totals != 0)
923
+ confusion_scores = 1.0 - torch.tensor(acc, device=device, dtype=torch.float32)
924
+ sigma = 0.5
925
+ gaussian = torch.exp(-((1 - class_similarity) ** 2) / (2 * sigma ** 2))
926
+ propagated = gaussian @ confusion_scores
927
+ v_sum = torch.zeros(5, device=device); v_cnt = torch.zeros(5, device=device)
928
+ for cls, v in enumerate(vertex_map.tolist()):
929
+ v_sum[v] += propagated[cls]; v_cnt[v] += 1
930
+ v_avg = torch.zeros_like(v_sum); mask = v_cnt > 0; v_avg[mask] = v_sum[mask] / v_cnt[mask]
931
+ vw = 1.0 - torch.tanh(v_avg)
932
+ return F.normalize(vw, p=1, dim=0) * 1.0
933
+
934
+ # ---------------------------------------------------------------------
935
+ # Evaluate / Train
936
+ # ---------------------------------------------------------------------
937
+ def evaluate(encoder: nn.Module, constellation: nn.Module, loader, num_classes: int, device: torch.device, collect_diag: bool = False):
938
+ encoder.eval(); constellation.eval()
939
+ all_preds, all_targets = [], []
940
+ lambda_vals = []
941
+ soft_w_sums = torch.zeros(5, device=device)
942
+ soft_w_count = 0
943
+ with torch.no_grad():
944
+ for x, y in tqdm(loader, desc="Evaluating"):
945
+ x, y = x.to(device), y.to(device)
946
+ if collect_diag:
947
+ z, diag = encoder(x, return_diag=True)
948
+ w = diag["softmax_weights"]
949
+ soft_w_sums += w.sum(dim=0); soft_w_count += w.size(0)
950
+ else:
951
+ z = encoder(x)
952
+ logits, _ = constellation(z)
953
+ preds = logits.argmax(dim=1)
954
+ all_preds.append(preds.cpu().numpy())
955
+ all_targets.append(y.cpu().numpy())
956
+ if hasattr(encoder, "opinion") and hasattr(encoder.opinion, "_lambda_raw"):
957
+ lambda_vals.append(float(torch.sigmoid(encoder.opinion._lambda_raw).item()))
958
+ all_preds = np.concatenate(all_preds); all_targets = np.concatenate(all_targets)
959
+ acc = float((all_preds == all_targets).mean())
960
+ cm = confusion_matrix(all_targets, all_preds, labels=np.arange(num_classes))
961
+ per_class = np.divide(cm.diagonal(), cm.sum(axis=1), out=np.zeros(num_classes), where=cm.sum(axis=1)!=0)
962
+ avg_soft_w = (soft_w_sums / soft_w_count).detach().cpu().numpy() if (collect_diag and soft_w_count > 0) else None
963
+ lam_eval = float(np.mean(lambda_vals)) if lambda_vals else None
964
+ return acc, per_class.tolist(), cm, avg_soft_w, lam_eval
965
+
966
+ def _adapt_pairs_by_classes(cfg: Dict, num_classes: int) -> int:
967
+ # Keep ~<=10 classes per vertex group across pairs
968
+ pairs = cfg.get("num_pentachoron_pairs", 1)
969
+ target = max(1, int(math.ceil(num_classes / 10)))
970
+ return max(pairs, target)
971
+
972
+ def train_one(
973
+ train_loader,
974
+ test_loader,
975
+ num_classes: int,
976
+ cfg: dict,
977
+ device: torch.device,
978
+ writer: SummaryWriter,
979
+ class_names: Optional[list] = None,
980
+ ):
981
+ pairs = _adapt_pairs_by_classes(cfg, num_classes)
982
+ if pairs != cfg.get("num_pentachoron_pairs"):
983
+ print(f"[auto] Adjusting num_pentachoron_pairs -> {pairs} for {num_classes} classes.")
984
+ cfg_local = dict(cfg); cfg_local["num_pentachoron_pairs"] = pairs
985
+
986
+ encoder = PentaFreqEncoderV2(
987
+ input_dim=cfg_local["input_dim"],
988
+ input_ch=cfg_local.get("input_channels", 1),
989
+ base_dim=cfg_local["base_dim"],
990
+ proj_dim=None,
991
+ num_heads=cfg_local.get("num_heads", 14),
992
+ channels=cfg_local.get("channels", 12),
993
+ ).to(device)
994
+
995
+ constellation = BatchedPentachoronConstellation(
996
+ num_classes=num_classes,
997
+ dim=cfg_local["base_dim"],
998
+ num_pairs=cfg_local["num_pentachoron_pairs"],
999
+ num_heads=cfg_local.get("num_heads", 14),
1000
+ device=device,
1001
+ lambda_sep=cfg_local["lambda_separation"],
1002
+ ).to(device)
1003
+
1004
+ diag_head = RoseDiagnosticHead(cfg_local["base_dim"]).to(device)
1005
+
1006
+ params = list(encoder.parameters()) + list(constellation.parameters()) + list(diag_head.parameters())
1007
+ optim = torch.optim.AdamW(params, lr=cfg_local["lr"], weight_decay=cfg_local["weight_decay"])
1008
+ lr_sched = torch.optim.lr_scheduler.CosineAnnealingLR(optim, T_max=cfg_local["epochs"])
1009
+
1010
+ w_ce = float(cfg_local.get("w_ce", 1.0))
1011
+ w_dual = float(cfg_local.get("w_dual", 1.0))
1012
+ w_rose = float(cfg_local.get("w_rose", 1.0))
1013
+ w_diag = float(cfg_local.get("w_diag", 0.1))
1014
+ w_reg = float(cfg_local.get("w_reg", cfg_local["loss_weight_scalar"]))
1015
+
1016
+ history = {"train_loss": [], "train_acc": [], "test_acc": [], "ce": [], "dual": [], "rose": [], "diag": [], "reg": [], "lambda": []}
1017
+ best = {"acc": 0.0, "cm": None, "epoch": -1}
1018
+ vertex_weights = None
1019
+
1020
+ global_step = 0
1021
+ for epoch in range(cfg_local["epochs"]):
1022
+ encoder.train(); constellation.train(); diag_head.train()
1023
+ sum_loss = sum_ce = sum_dual = sum_rose = sum_diag = sum_reg = 0.0
1024
+ correct = total = 0
1025
+
1026
+ pbar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{cfg_local['epochs']} [Train]")
1027
+ for x, y in pbar:
1028
+ x, y = x.to(device), y.to(device)
1029
+ optim.zero_grad()
1030
+
1031
+ z = encoder(x)
1032
+ logits, _ = constellation(z)
1033
+
1034
+ l_ce = F.cross_entropy(logits, y)
1035
+ l_dual = dual_contrastive_loss(z, y, constellation, temp=cfg_local["temp"])
1036
+ # Around line 1000 in train_one, replace:
1037
+ #l_rose, rose_scores = rose_contrastive_loss(z, y, constellation, temp=cfg_local["temp"])
1038
+ #pred_rose = diag_head(z.detach()).squeeze(-1)
1039
+ #l_diag = F.mse_loss(pred_rose, rose_scores)
1040
+
1041
+ # With:
1042
+ if w_rose > 0 or w_diag > 0:
1043
+ l_rose, rose_scores = rose_contrastive_loss(z, y, constellation, temp=cfg_local["temp"])
1044
+ pred_rose = diag_head(z.detach()).squeeze(-1)
1045
+ l_diag = F.mse_loss(pred_rose, rose_scores)
1046
+ else:
1047
+ l_rose = torch.tensor(0.0, device=device)
1048
+ l_diag = torch.tensor(0.0, device=device)
1049
+ rose_scores = None
1050
+ #pred_rose = diag_head(z.detach()).squeeze(-1)
1051
+ #l_diag = F.mse_loss(pred_rose, rose_scores)
1052
+ l_reg = constellation.regularization_loss(vertex_weights=vertex_weights)
1053
+
1054
+ loss = (w_ce*l_ce) + (w_dual*l_dual) + (w_rose*l_rose) + (w_diag*l_diag) + (w_reg*l_reg)
1055
+
1056
+ # after computing l_ce, l_dual, l_rose, l_diag, l_reg and loss
1057
+ if not torch.isfinite(l_ce) or not torch.isfinite(l_dual) \
1058
+ or not torch.isfinite(l_rose) or not torch.isfinite(l_diag) \
1059
+ or not torch.isfinite(l_reg) or not torch.isfinite(loss):
1060
+ print("[NaN-guard] non-finite detected. Skipping step. "
1061
+ f"ce={l_ce.item() if torch.isfinite(l_ce) else 'nan'}, "
1062
+ f"dual={l_dual.item() if torch.isfinite(l_dual) else 'nan'}, "
1063
+ f"rose={l_rose.item() if torch.isfinite(l_rose) else 'nan'}, "
1064
+ f"reg={l_reg.item() if torch.isfinite(l_reg) else 'nan'}, "
1065
+ )
1066
+
1067
+ # Soft defuse: drop LR a bit for stability
1068
+ for g in optim.param_groups:
1069
+ g["lr"] = max(g["lr"] * 0.5, 1e-6)
1070
+
1071
+ optim.zero_grad(set_to_none=True)
1072
+ # Optional: clip parameter norms right now to kill accidental blow-ups
1073
+ with torch.no_grad():
1074
+ for p in list(encoder.parameters()) + list(constellation.parameters()):
1075
+ if torch.isfinite(p).all():
1076
+ p.clamp_(-1e3, 1e3)
1077
+ continue
1078
+ loss.backward()
1079
+ torch.nn.utils.clip_grad_norm_(encoder.parameters(), 1.0)
1080
+ torch.nn.utils.clip_grad_norm_(constellation.parameters(), 1.0)
1081
+ optim.step()
1082
+
1083
+ bs = x.size(0)
1084
+ sum_loss += loss.item() * bs
1085
+ sum_ce += l_ce.item() * bs
1086
+ sum_dual += l_dual.item() * bs
1087
+ sum_rose += l_rose.item() * bs
1088
+ sum_diag += l_diag.item() * bs
1089
+ sum_reg += l_reg.item() * bs
1090
+
1091
+ preds = logits.argmax(dim=1)
1092
+ correct += (preds == y).sum().item()
1093
+ total += bs
1094
+
1095
+ # TB (step)
1096
+ writer.add_scalar("step/loss", loss.item(), global_step)
1097
+ writer.add_scalar("step/ce", l_ce.item(), global_step)
1098
+ writer.add_scalar("step/dual", l_dual.item(), global_step)
1099
+ writer.add_scalar("step/rose", l_rose.item(), global_step)
1100
+ writer.add_scalar("step/diag", l_diag.item(), global_step)
1101
+ writer.add_scalar("step/reg", l_reg.item(), global_step)
1102
+ global_step += 1
1103
+
1104
+ pbar.set_postfix({
1105
+ "loss": f"{loss.item():.4f}",
1106
+ "acc": f"{correct/max(1,total):.4f}",
1107
+ "ce": f"{l_ce.item():.3f}",
1108
+ "dual": f"{l_dual.item():.3f}",
1109
+ "rose": f"{l_rose.item():.3f}",
1110
+ "reg": f"{l_reg.item():.3f}",
1111
+ "blend": f"{constellation.blend_weight.item()}",
1112
+ })
1113
+
1114
+ train_loss = sum_loss / max(1, total)
1115
+ train_acc = correct / max(1, total)
1116
+ history["train_loss"].append(train_loss)
1117
+ history["train_acc"].append(train_acc)
1118
+ history["ce"].append(sum_ce / max(1,total))
1119
+ history["dual"].append(sum_dual / max(1,total))
1120
+ history["rose"].append(sum_rose / max(1,total))
1121
+ history["diag"].append(sum_diag / max(1,total))
1122
+ history["reg"].append(sum_reg / max(1,total))
1123
+
1124
+ # Eval
1125
+ test_acc, per_class_acc, cm, avg_soft_w, lam_eval = evaluate(
1126
+ encoder, constellation, test_loader, num_classes, device, collect_diag=True
1127
+ )
1128
+ history["test_acc"].append(test_acc)
1129
+ if lam_eval is not None:
1130
+ history["lambda"].append(lam_eval)
1131
+ else:
1132
+ history["lambda"].append(float(torch.sigmoid(encoder.opinion._lambda_raw).item())
1133
+ if hasattr(encoder, "opinion") else 0.5)
1134
+
1135
+ lr_sched.step()
1136
+
1137
+ # TB (epoch)
1138
+ writer.add_scalar("epoch/train_loss", train_loss, epoch+1)
1139
+ writer.add_scalar("epoch/train_acc", train_acc, epoch+1)
1140
+ writer.add_scalar("epoch/test_acc", test_acc, epoch+1)
1141
+ writer.add_scalar("epoch/lr", optim.param_groups[0]["lr"], epoch+1)
1142
+ writer.add_scalar("epoch/lambda", history["lambda"][-1], epoch+1)
1143
+
1144
+ print(f"\n[Epoch {epoch+1}/{cfg_local['epochs']}] "
1145
+ f"TrainLoss {train_loss:.4f} | TrainAcc {train_acc:.4f} | TestAcc {test_acc:.4f} | "
1146
+ f"CE {history['ce'][-1]:.3f} Dual {history['dual'][-1]:.3f} "
1147
+ f"ROSE {history['rose'][-1]:.3f} Reg {history['reg'][-1]:.3f} λ {history['lambda'][-1]:.3f}")
1148
+
1149
+ # Update reg weights
1150
+ with torch.no_grad():
1151
+ class_sim = get_class_similarity(constellation, num_classes).to(device)
1152
+ vertex_weights = vertex_weights_from_confusion(cm, class_sim, constellation.vertex_map, device)
1153
+
1154
+ if test_acc > best["acc"]:
1155
+ best["acc"], best["cm"], best["epoch"] = test_acc, cm, epoch+1
1156
+ print(f" 🎯 New Best Acc: {best['acc']:.4f} at epoch {best['epoch']}")
1157
+
1158
+ # Optional confusion per epoch
1159
+ try:
1160
+ import matplotlib.pyplot as plt
1161
+ import seaborn as sns
1162
+ os.makedirs("plots", exist_ok=True)
1163
+ plt.figure(figsize=(9, 7))
1164
+ sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
1165
+ xticklabels=class_names, yticklabels=class_names)
1166
+ plt.title(f'Confusion (Epoch {epoch+1}) | Acc: {test_acc:.4f}')
1167
+ plt.xlabel('Predicted'); plt.ylabel('True'); plt.tight_layout()
1168
+ plt.savefig(f'plots/confusion_epoch_{epoch+1}.png', dpi=150)
1169
+ plt.close()
1170
+ except Exception as e:
1171
+ print(f"(Confusion heatmap skipped this epoch: {e})")
1172
+
1173
+ return encoder, constellation, diag_head, history, best
1174
+
1175
+ # ---------------------------------------------------------------------
1176
+ # Plots (local convenience; artifacts already keep TB)
1177
+ # ---------------------------------------------------------------------
1178
+ def plot_history(history: dict, outdir: str = "plots"):
1179
+ os.makedirs(outdir, exist_ok=True)
1180
+ import matplotlib.pyplot as plt
1181
+ plt.figure(figsize=(10,5))
1182
+ plt.plot(history['train_acc'], label='Train Acc')
1183
+ plt.plot(history['test_acc'], label='Test Acc')
1184
+ plt.title('Accuracy over Epochs'); plt.xlabel('Epoch'); plt.ylabel('Accuracy'); plt.legend(); plt.grid(True, ls='--', alpha=0.4)
1185
+ plt.tight_layout(); plt.savefig(f"{outdir}/accuracy.png", dpi=150); plt.close()
1186
+
1187
+ plt.figure(figsize=(10,5))
1188
+ plt.plot(history['train_loss'], label='Total')
1189
+ plt.plot(history['ce'], label='CE')
1190
+ plt.plot(history['dual'], label='DualNCE')
1191
+ plt.plot(history['rose'], label='ROSE')
1192
+ plt.plot(history['reg'], label='Reg')
1193
+ plt.title('Loss Components'); plt.xlabel('Epoch'); plt.ylabel('Loss'); plt.legend(); plt.grid(True, ls='--', alpha=0.4)
1194
+ plt.tight_layout(); plt.savefig(f"{outdir}/loss_components.png", dpi=150); plt.close()
1195
+
1196
+ plt.figure(figsize=(8,4))
1197
+ plt.plot(history['lambda'])
1198
+ plt.title('λ (Geometry ↔ Attention Gate)'); plt.xlabel('Epoch'); plt.ylabel('λ'); plt.grid(True, ls='--', alpha=0.4)
1199
+ plt.tight_layout(); plt.savefig(f"{outdir}/lambda.png", dpi=150); plt.close()
1200
+
1201
+ def plot_confusion(cm: np.ndarray, class_names: list, outpath: str):
1202
+ import matplotlib.pyplot as plt
1203
+ try:
1204
+ import seaborn as sns
1205
+ plt.figure(figsize=(10,8))
1206
+ sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
1207
+ xticklabels=class_names, yticklabels=class_names)
1208
+ plt.title('Best Confusion Matrix'); plt.xlabel('Predicted'); plt.ylabel('True')
1209
+ plt.tight_layout(); plt.savefig(outpath, dpi=150); plt.close()
1210
+ except Exception:
1211
+ plt.figure(figsize=(10,8))
1212
+ plt.imshow(cm, aspect='auto'); plt.title('Best Confusion Matrix')
1213
+ plt.xlabel('Predicted'); plt.ylabel('True'); plt.colorbar()
1214
+ plt.tight_layout(); plt.savefig(outpath, dpi=150); plt.close()
1215
+
1216
+
1217
+ def _sanitize_for_path(s: str, preserve_case: bool = True) -> str:
1218
+ """Keep letters/digits/._- ; replace others with '_'."""
1219
+ out = []
1220
+ for ch in (s if preserve_case else s.lower()):
1221
+ if ch.isalnum() or ch in "._-":
1222
+ out.append(ch)
1223
+ else:
1224
+ out.append("_")
1225
+ return "".join(out)
1226
+
1227
+ def _build_hf_paths(dataset_name: str, cfg: Dict, ts: str) -> Dict[str, str]:
1228
+ pres = bool(cfg.get("hf_preserve_case", True))
1229
+ dataset_disp = dataset_name if pres else dataset_name.lower()
1230
+ dataset_token = _sanitize_for_path(dataset_disp, preserve_case=pres)
1231
+ slug = _dataset_slug(dataset_name) # lowercase with '+' for sweeps (kept for convenience)
1232
+
1233
+ root = cfg.get("hf_subdir_root", "pentachora-adaptive-encoded").strip("/")
1234
+
1235
+ # templates allow {dataset}, {slug}, {ts}
1236
+ dtempl = cfg.get("hf_dataset_dir_template", "{dataset}")
1237
+ rtempl = cfg.get("hf_run_dir_template", "{ts}_{dataset}")
1238
+
1239
+ dataset_dir = dtempl.format(dataset=dataset_token, slug=slug, ts=ts)
1240
+ run_dir = rtempl.format(dataset=dataset_token, slug=slug, ts=ts)
1241
+
1242
+ path_in_repo = f"{root}/{dataset_dir}/{run_dir}".strip("/")
1243
+ local_root = Path("artifacts") / root / dataset_dir / run_dir
1244
+
1245
+ return {
1246
+ "dataset_token": dataset_token,
1247
+ "path_in_repo": path_in_repo,
1248
+ "local_root": str(local_root),
1249
+ }
1250
+
1251
+
1252
+
1253
+ def save_and_push_artifacts(
1254
+ *,
1255
+ encoder: nn.Module,
1256
+ constellation: nn.Module,
1257
+ diag_head: nn.Module,
1258
+ config: Dict,
1259
+ class_names: List[str],
1260
+ history: Dict,
1261
+ best: Dict,
1262
+ tb_log_dir: Path,
1263
+ dataset_names: List[str], # pass a single dataset name here
1264
+ ):
1265
+ assert len(dataset_names) == 1, "Pass a single dataset name to save_and_push_artifacts"
1266
+ dataset_name = dataset_names[0]
1267
+
1268
+ ts = _timestamp()
1269
+ repo_id = _resolve_repo_id(config)
1270
+ _hf_login_if_needed()
1271
+ api = _ensure_repo(repo_id)
1272
+
1273
+ paths = _build_hf_paths(dataset_name, config, ts)
1274
+ base_out = Path(paths["local_root"])
1275
+ base_out.mkdir(parents=True, exist_ok=True)
1276
+
1277
+ # Weight file naming
1278
+ ds_token = paths["dataset_token"]
1279
+ suffix = f"_{ds_token}" if bool(config.get("hf_weight_suffix_dataset", True)) else ""
1280
+
1281
+ wdir = base_out / "weights"; wdir.mkdir(parents=True, exist_ok=True)
1282
+ save_safetensors({k: v.cpu() for k, v in encoder.state_dict().items()}, str(wdir / f"encoder{suffix}.safetensors"))
1283
+ save_safetensors({k: v.cpu() for k, v in constellation.state_dict().items()}, str(wdir / f"constellation{suffix}.safetensors"))
1284
+ save_safetensors({k: v.cpu() for k, v in diag_head.state_dict().items()}, str(wdir / f"diagnostic_head{suffix}.safetensors"))
1285
+
1286
+ # Config + history
1287
+ (base_out / "config.json").write_text(json.dumps(config, indent=2, sort_keys=True), encoding="utf-8")
1288
+ (base_out / "history.json").write_text(json.dumps(history, indent=2, sort_keys=True), encoding="utf-8")
1289
+
1290
+ # CSV history
1291
+ max_len = max(len(history.get("train_loss", [])), len(history.get("train_acc", [])), len(history.get("test_acc", [])))
1292
+ df = pd.DataFrame({
1293
+ "epoch": list(range(1, max_len + 1)),
1294
+ "train_loss": history.get("train_loss", [np.nan]*max_len),
1295
+ "train_acc": history.get("train_acc", [np.nan]*max_len),
1296
+ "test_acc": history.get("test_acc", [np.nan]*max_len),
1297
+ "ce": history.get("ce", [np.nan]*max_len),
1298
+ "dual": history.get("dual", [np.nan]*max_len),
1299
+ "rose": history.get("rose", [np.nan]*max_len),
1300
+ "diag": history.get("diag", [np.nan]*max_len),
1301
+ "reg": history.get("reg", [np.nan]*max_len),
1302
+ "lambda": history.get("lambda", [np.nan]*max_len),
1303
+ })
1304
+ df.to_csv(base_out / "history.csv", index=False)
1305
+
1306
+ # Plots
1307
+ if Path("plots").exists():
1308
+ shutil.copytree("plots", base_out / "plots", dirs_exist_ok=True)
1309
+
1310
+ # TensorBoard
1311
+ tb_dst = base_out / "tensorboard"; tb_dst.mkdir(parents=True, exist_ok=True)
1312
+ if tb_log_dir and Path(tb_log_dir).exists():
1313
+ for p in Path(tb_log_dir).glob("*"):
1314
+ shutil.copy2(p, tb_dst / p.name)
1315
+ _zip_dir(tb_dst, base_out / "tensorboard_events.zip")
1316
+
1317
+ # Manifest + README
1318
+ manifest = {
1319
+ "timestamp": ts,
1320
+ "repo_id": repo_id,
1321
+ "subdirectory": paths["path_in_repo"],
1322
+ "dataset_name": dataset_name,
1323
+ "class_names": class_names,
1324
+ "num_classes": len(class_names),
1325
+ "models": {
1326
+ "encoder": {"params": _param_count(encoder)},
1327
+ "constellation": {"params": _param_count(constellation)},
1328
+ "diagnostic_head": {"params": _param_count(diag_head)},
1329
+ },
1330
+ "results": {
1331
+ "best_test_accuracy": float(best.get("acc", 0.0)),
1332
+ "best_epoch": int(best.get("epoch", -1)),
1333
+ },
1334
+ "environment": {
1335
+ "python": sys.version,
1336
+ "platform": platform.platform(),
1337
+ "torch": torch.__version__,
1338
+ "cuda_available": torch.cuda.is_available(),
1339
+ "cuda_device": (torch.cuda.get_device_name(0) if torch.cuda.is_available() else None),
1340
+ "cpu_count": psutil.cpu_count(logical=True),
1341
+ "memory_gb": round(psutil.virtual_memory().total / (1024**3), 2),
1342
+ },
1343
+ }
1344
+ (base_out / "manifest.json").write_text(json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8")
1345
+
1346
+ (base_out / "README.md").write_text(
1347
+ f"""# Pentachora Adaptive Encoded — {ts}
1348
+
1349
+ **Dataset:** {dataset_name}
1350
+
1351
+ **Contents**
1352
+ - `weights/*.safetensors` — encoder, constellation, diagnostic head
1353
+ - `config.json`, `manifest.json`
1354
+ - `history.json` / `history.csv`
1355
+ - `tensorboard/` (and `tensorboard_events.zip`)
1356
+ - `plots/` — accuracy, loss, λ, confusion
1357
+ """,
1358
+ encoding="utf-8"
1359
+ )
1360
+
1361
+ # Push
1362
+ print(f"[push] Uploading to hf://{repo_id}/{paths['path_in_repo']}")
1363
+ api.upload_folder(
1364
+ repo_id=repo_id,
1365
+ folder_path=str(base_out),
1366
+ path_in_repo=paths["path_in_repo"],
1367
+ repo_type="model",
1368
+ commit_message=f"[{dataset_name}] {ts} | best_acc={manifest['results']['best_test_accuracy']:.4f}",
1369
+ )
1370
+ print("[push] ✅ Upload complete.")
1371
+ return base_out
1372
+
1373
+ # ---------------------------------------------------------------------
1374
+ # Dataset sweep
1375
+ # ---------------------------------------------------------------------
1376
+ def run_one_dataset(name: str) -> Dict:
1377
+ print("\n" + "="*60)
1378
+ print(f"RUN: {name}")
1379
+ print("="*60)
1380
+
1381
+ # Load
1382
+ train_loader, test_loader, ncls, class_names, in_dim, out_c = get_dataset_single(
1383
+ name, batch_size=config["batch_size"], num_workers=2
1384
+ )
1385
+ cfg_local = dict(config)
1386
+ cfg_local["num_classes"] = ncls
1387
+ cfg_local["input_dim"] = in_dim
1388
+ cfg_local["input_channels"] = out_c
1389
+
1390
+ # TB writer per dataset
1391
+ ts = _timestamp()
1392
+ tb_dir = Path("tb_logs") / f"{_dataset_slug(name)}" / ts
1393
+ tb_dir.mkdir(parents=True, exist_ok=True)
1394
+ writer = SummaryWriter(log_dir=str(tb_dir))
1395
+
1396
+ start = time.time()
1397
+ encoder, constellation, diag_head, history, best = train_one(
1398
+ train_loader, test_loader, ncls, cfg_local, device, writer, class_names
1399
+ )
1400
+ elapsed_min = (time.time() - start) / 60.0
1401
+
1402
+ # Plots
1403
+ plot_history(history, outdir="plots")
1404
+ if best["cm"] is not None:
1405
+ plot_confusion(best["cm"], class_names, outpath=f"plots/best_confusion_{_dataset_slug(name)}_epoch_{best['epoch']}.png")
1406
+
1407
+ # Push artifacts
1408
+ save_and_push_artifacts(
1409
+ encoder=encoder,
1410
+ constellation=constellation,
1411
+ diag_head=diag_head,
1412
+ config=cfg_local,
1413
+ class_names=class_names,
1414
+ history=history,
1415
+ best=best,
1416
+ tb_log_dir=tb_dir,
1417
+ dataset_names=[name],
1418
+ )
1419
+
1420
+ result = {
1421
+ "dataset": name,
1422
+ "classes": ncls,
1423
+ "channels": out_c,
1424
+ "img_size": config.get("img_size", 28),
1425
+ "best_acc": float(best["acc"]),
1426
+ "best_epoch": int(best["epoch"]),
1427
+ "params_encoder": _param_count(encoder),
1428
+ "params_constellation": _param_count(constellation),
1429
+ "elapsed_min": round(elapsed_min, 3),
1430
+ }
1431
+ print(f"[done] {name} -> best_acc={result['best_acc']:.4f} @ epoch {result['best_epoch']} time={result['elapsed_min']:.2f}m")
1432
+ return result
1433
+
1434
+ def run_sweep(datasets: List[str]) -> Dict:
1435
+ os.makedirs("sweeps", exist_ok=True)
1436
+ results = []
1437
+ failures = []
1438
+ for name in datasets:
1439
+ try:
1440
+ results.append(run_one_dataset(name))
1441
+ except Exception as e:
1442
+ print(f"[fail] {name}: {e}")
1443
+ failures.append({"dataset": name, "error": str(e)})
1444
+
1445
+ # Save local sweep summary
1446
+ ts = _timestamp()
1447
+ sweep_dir = Path("sweeps") / ts
1448
+ sweep_dir.mkdir(parents=True, exist_ok=True)
1449
+
1450
+ df = pd.DataFrame(results)
1451
+ df.to_csv(sweep_dir / "results.csv", index=False)
1452
+ (sweep_dir / "results.json").write_text(json.dumps(results, indent=2), encoding="utf-8")
1453
+ (sweep_dir / "failures.json").write_text(json.dumps(failures, indent=2), encoding="utf-8")
1454
+
1455
+ # Push sweep summary
1456
+ repo_id = _resolve_repo_id(config)
1457
+ _hf_login_if_needed()
1458
+ api = _ensure_repo(repo_id)
1459
+ path_in_repo = f"pentachora-adaptive-encoded/_sweep/{ts}"
1460
+ print(f"[push] Uploading sweep summary to hf://{repo_id}/{path_in_repo}")
1461
+ api.upload_folder(repo_id=repo_id, folder_path=str(sweep_dir), path_in_repo=path_in_repo, repo_type="model")
1462
+ print("[push] ✅ Sweep summary uploaded.")
1463
+
1464
+ return {"timestamp": ts, "results": results, "failures": failures, "path_in_repo": path_in_repo}
1465
+
1466
+ # ---------------------------------------------------------------------
1467
+ # Main
1468
+ # ---------------------------------------------------------------------
1469
+ def main():
1470
+ print("\n" + "="*60)
1471
+ print("PENTACHORON CONSTELLATION FINAL CONFIGURATION")
1472
+ print("="*60)
1473
+ for k, v in config.items():
1474
+ print(f"{k:24}: {v}")
1475
+ if config["lr"] > 1e-1:
1476
+ print(f"⚠️ High LR detected ({config['lr']}). If unstable, try 5e-3 to 5e-2.")
1477
+
1478
+ # Sweep mode?
1479
+ if bool(config.get("sweep_all", False)) or os.getenv("RUN_SWEEP", "0") == "1":
1480
+ # Try all TorchVision + MedMNIST (skip those not available)
1481
+ datasets_all = list(TORCHVISION_DATASETS)
1482
+ if medmnist is not None:
1483
+ datasets_all += MEDMNIST_DATASETS
1484
+ out = run_sweep(datasets_all)
1485
+ print(f"\nSweep complete. Summary path: {out['path_in_repo']}")
1486
+ return
1487
+
1488
+
1489
+
1490
+
1491
+ # Single dataset default (edit here as desired)
1492
+ DATASET = config.get("dataset", "FashionMNIST")
1493
+
1494
+
1495
+
1496
+
1497
+
1498
+ train_loader, test_loader, ncls, class_names, in_dim, out_c = get_dataset_single(
1499
+ DATASET, batch_size=config["batch_size"], num_workers=2
1500
+ )
1501
+ config["num_classes"] = ncls
1502
+ config["input_dim"] = in_dim
1503
+ config["input_channels"] = out_c
1504
+
1505
+ tb_dir = Path("tb_logs") / f"{_dataset_slug(DATASET)}" / _timestamp()
1506
+ tb_dir.mkdir(parents=True, exist_ok=True)
1507
+ writer = SummaryWriter(log_dir=str(tb_dir))
1508
+
1509
+ start = time.time()
1510
+ encoder, constellation, diag_head, history, best = train_one(
1511
+ train_loader, test_loader, ncls, config, device, writer, class_names
1512
+ )
1513
+ elapsed = (time.time() - start) / 60.0
1514
+
1515
+ # Plots
1516
+ plot_history(history, outdir="plots")
1517
+ if best["cm"] is not None:
1518
+ plot_confusion(best["cm"], class_names, outpath=f"plots/best_confusion_epoch_{best['epoch']}.png")
1519
+
1520
+ print("\n" + "="*60)
1521
+ print("TRAINING COMPLETE")
1522
+ print("="*60)
1523
+ print(f"Best Test Accuracy : {best['acc']*100:.2f}% @ epoch {best['epoch']}")
1524
+ print(f"Total Training Time: {elapsed:.2f} minutes")
1525
+
1526
+ save_and_push_artifacts(
1527
+ encoder=encoder,
1528
+ constellation=constellation,
1529
+ diag_head=diag_head,
1530
+ config=config,
1531
+ class_names=class_names,
1532
+ history=history,
1533
+ best=best,
1534
+ tb_log_dir=tb_dir,
1535
+ dataset_names=[DATASET],
1536
+ )
1537
+ print("[done] Artifacts uploaded and saved locally.")
1538
+
1539
+ if __name__ == "__main__":
1540
+ main()