| """ |
| 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__) |
|
|
|
|
| |
|
|
| 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: |
| |
| 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 |
| 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 |
|
|
|
|
| |
|
|
| 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 |
|
|
|
|
| |
|
|
| 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] |
|
|
| |
| if self.noise_scale > 0: |
| x = x + torch.randn_like(x) * self.noise_scale |
|
|
| |
| 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 |
|
|
|
|
| |
|
|
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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)") |
|
|
| |
| X, Y, cids = extract_arrays(all_samples, input_features, output_stats) |
| logger.info(f"Arrays: X={X.shape}, Y={Y.shape}") |
|
|
| |
| 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)" |
| ) |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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, |
| 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 |
|
|