""" PyTorch Dataset for loading circuit simulation JSONs. Each circuit file (circuit_XXXXX.json) contains a list of 11 dicts — one per ACh level. Each dict has circuit config fields + a nested "statistics" dict with the 11 summary stats. Two dataset variants: - SimDatasetB: all ACh levels (for Model B, ~55K samples) - SimDatasetA: only ACh=0.0 rows (for Model A, ~5K samples) Normalization is fitted on training data and applied consistently. """ from __future__ import annotations import glob import json import logging import math from pathlib import Path import numpy as np import torch from torch.utils.data import Dataset from .config import ( INPUT_FEATURES_A, INPUT_FEATURES_B, LOG_TRANSFORM_INPUTS, LOG_TRANSFORM_STATS, OUTPUT_STATS, TrainConfig, ) logger = logging.getLogger(__name__) # ── Normalization helpers ──────────────────────────────────────────────────── class Normalizer: """Z-score normalizer that can optionally log-transform columns first. Usage: norm = Normalizer(log_cols={2, 5}) norm.fit(X_train) # compute mean/std from training data X_normed = norm.transform(X) X_orig = norm.inverse(X_normed) """ def __init__(self, log_cols: set[int] | None = None, eps: float = 1e-8): self.log_cols = log_cols or set() self.eps = eps self.mean: np.ndarray | None = None self.std: np.ndarray | None = None def _log_transform(self, X: np.ndarray) -> np.ndarray: X = X.copy() for c in self.log_cols: # Signed log1p: preserves sign, handles negatives X[:, c] = np.sign(X[:, c]) * np.log1p(np.abs(X[:, c])) return X def _log_inverse(self, X: np.ndarray) -> np.ndarray: X = X.copy() for c in self.log_cols: X[:, c] = np.sign(X[:, c]) * np.expm1(np.abs(X[:, c])) return X def fit(self, X: np.ndarray) -> "Normalizer": X_t = self._log_transform(X) self.mean = X_t.mean(axis=0) self.std = X_t.std(axis=0) self.std[self.std < self.eps] = 1.0 # avoid /0 for constant cols return self def transform(self, X: np.ndarray) -> np.ndarray: assert self.mean is not None, "Call .fit() first" return (self._log_transform(X) - self.mean) / self.std def inverse(self, X: np.ndarray) -> np.ndarray: assert self.mean is not None, "Call .fit() first" return self._log_inverse(X * self.std + self.mean) def state_dict(self) -> dict: return { "log_cols": sorted(self.log_cols), "eps": self.eps, "mean": self.mean.tolist() if self.mean is not None else None, "std": self.std.tolist() if self.std is not None else None, } @classmethod def from_state_dict(cls, d: dict) -> "Normalizer": n = cls(log_cols=set(d["log_cols"]), eps=d["eps"]) if d["mean"] is not None: n.mean = np.array(d["mean"]) n.std = np.array(d["std"]) return n # ── Raw data loading ───────────────────────────────────────────────────────── def load_circuit_jsons(sim_dir: str, extra_dirs: list[str] | None = None) -> list[dict]: """Load all circuit JSON files and flatten into a list of sample dicts. Each file has 11 entries (one per ACh level) or 1 entry (ACh=0 only). We return a flat list of all samples across all circuits. Args: sim_dir: Primary simulation directory. extra_dirs: Additional directories to load from (e.g., ach0_extra). """ all_dirs = [sim_dir] + (extra_dirs or []) samples: list[dict] = [] total_files = 0 for d in all_dirs: pattern = str(Path(d) / "circuit_*.json") files = sorted(glob.glob(pattern)) if not files: logger.warning(f"No circuit files found in {d} — skipping") continue logger.info(f"Found {len(files)} circuit files in {d}") total_files += len(files) for fpath in files: with open(fpath) as f: circuit_data = json.load(f) if isinstance(circuit_data, list): samples.extend(circuit_data) else: samples.append(circuit_data) if not samples: raise FileNotFoundError(f"No circuit files found in any of: {all_dirs}") logger.info(f"Loaded {len(samples)} total samples from {total_files} files across {len(all_dirs)} directories") return samples def extract_arrays( samples: list[dict], input_features: list[str], output_stats: list[str], ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Convert list of sample dicts → (X_inputs, Y_targets, circuit_ids). Returns: X: (N, n_input_features) float64 Y: (N, n_output_stats) float64 cids: (N,) int64 — circuit IDs for splitting """ N = len(samples) n_in = len(input_features) n_out = len(output_stats) X = np.zeros((N, n_in), dtype=np.float64) Y = np.zeros((N, n_out), dtype=np.float64) cids = np.zeros(N, dtype=np.int64) for i, s in enumerate(samples): cids[i] = s["circuit_id"] for j, feat in enumerate(input_features): X[i, j] = float(s[feat]) stats = s["statistics"] for j, stat in enumerate(output_stats): Y[i, j] = float(stats[stat]) return X, Y, cids # ── Dataset classes ────────────────────────────────────────────────────────── class SimDataset(Dataset): """PyTorch Dataset wrapping normalized input/output tensors.""" def __init__(self, X: torch.Tensor, Y: torch.Tensor, noise_scale: float = 0.0, mixup_alpha: float = 0.0): assert X.shape[0] == Y.shape[0] self.X = X self.Y = Y self.noise_scale = noise_scale self.mixup_alpha = mixup_alpha def __len__(self) -> int: return self.X.shape[0] def __getitem__(self, idx: int) -> tuple[torch.Tensor, torch.Tensor]: x, y = self.X[idx], self.Y[idx] # Input noise augmentation (training only — caller sets noise_scale=0 for val) if self.noise_scale > 0: x = x + torch.randn_like(x) * self.noise_scale # Mixup augmentation if self.mixup_alpha > 0 and self.training_mode: lam = np.random.beta(self.mixup_alpha, self.mixup_alpha) j = np.random.randint(0, len(self.X)) x = lam * x + (1 - lam) * self.X[j] y = lam * y + (1 - lam) * self.Y[j] return x, y @property def training_mode(self) -> bool: return self.noise_scale > 0 or self.mixup_alpha > 0 # ── Builder functions ──────────────────────────────────────────────────────── def _log_col_indices(feature_names: list[str], log_set: set[str]) -> set[int]: """Find column indices that need log-transform.""" return {i for i, name in enumerate(feature_names) if name in log_set} def build_datasets( cfg: TrainConfig, model_variant: str = "B", ) -> tuple[SimDataset, SimDataset, Normalizer, Normalizer, dict]: """Load data, split, normalize, return (train_ds, val_ds, x_norm, y_norm, meta). Args: cfg: Training configuration. model_variant: "A" for plain HH (ACh=0 only), "B" for HH+ACh (all levels). Returns: train_ds: Training SimDataset val_ds: Validation SimDataset x_norm: Fitted Normalizer for inputs y_norm: Fitted Normalizer for outputs meta: Dict with split info, feature names, etc. """ assert model_variant in ("A", "B"), f"Unknown variant: {model_variant}" input_features = INPUT_FEATURES_A if model_variant == "A" else INPUT_FEATURES_B output_stats = OUTPUT_STATS # 1. Load all samples (from primary + extra directories) import os extra_dirs = [] if hasattr(cfg, "extra_ach0_dir") and cfg.extra_ach0_dir and os.path.isdir(cfg.extra_ach0_dir): extra_dirs.append(cfg.extra_ach0_dir) all_samples = load_circuit_jsons(cfg.sim_dir, extra_dirs=extra_dirs if extra_dirs else None) # 2. Filter for Model A (ACh=0 only) if model_variant == "A": all_samples = [s for s in all_samples if abs(s["ach_level"]) < 1e-6] logger.info(f"Model A: filtered to {len(all_samples)} samples (ACh=0 only)") # 3. Extract arrays X, Y, cids = extract_arrays(all_samples, input_features, output_stats) logger.info(f"Arrays: X={X.shape}, Y={Y.shape}") # 4. Train/val split BY CIRCUIT ID (prevents data leakage) unique_cids = np.unique(cids) rng = np.random.RandomState(cfg.seed) rng.shuffle(unique_cids) n_val = max(1, int(len(unique_cids) * cfg.val_frac)) val_cids = set(unique_cids[:n_val].tolist()) train_cids = set(unique_cids[n_val:].tolist()) train_mask = np.array([c in train_cids for c in cids]) val_mask = ~train_mask X_train, Y_train = X[train_mask], Y[train_mask] X_val, Y_val = X[val_mask], Y[val_mask] logger.info( f"Split: {len(train_cids)} train circuits ({X_train.shape[0]} samples), " f"{len(val_cids)} val circuits ({X_val.shape[0]} samples)" ) # 5. Fit normalizers on training data x_log_cols = _log_col_indices(input_features, LOG_TRANSFORM_INPUTS) y_log_cols = _log_col_indices(output_stats, LOG_TRANSFORM_STATS) x_norm = Normalizer(log_cols=x_log_cols).fit(X_train) y_norm = Normalizer(log_cols=y_log_cols).fit(Y_train) # 6. Transform X_train_n = x_norm.transform(X_train) X_val_n = x_norm.transform(X_val) Y_train_n = y_norm.transform(Y_train) Y_val_n = y_norm.transform(Y_val) # 7. To tensors (with augmentation for training set) train_ds = SimDataset( torch.tensor(X_train_n, dtype=torch.float32), torch.tensor(Y_train_n, dtype=torch.float32), noise_scale=cfg.aug_noise_scale, mixup_alpha=cfg.aug_mixup_alpha, ) val_ds = SimDataset( torch.tensor(X_val_n, dtype=torch.float32), torch.tensor(Y_val_n, dtype=torch.float32), noise_scale=0.0, # No augmentation on validation mixup_alpha=0.0, ) meta = { "model_variant": model_variant, "input_features": input_features, "output_stats": output_stats, "n_train_circuits": len(train_cids), "n_val_circuits": len(val_cids), "n_train_samples": int(X_train.shape[0]), "n_val_samples": int(X_val.shape[0]), "n_total_samples": int(X.shape[0]), } return train_ds, val_ds, x_norm, y_norm, meta