#!/usr/bin/env python3 """Merge per-(fold, seed) norm_stats JSON files into per-fold aggregates. Plan 5 Task 1B.3 — input is a directory of `norm_stats_foldN_seedM.json` files written by training. Each file contains per-feature normalisation statistics (mean, std) computed on that (fold, seed)'s training window. Across seeds for the same fold, those stats should be near-identical (same training window), but seed-driven sampling order can introduce small variation; this script collapses N seeds into one per-fold stat for the `evaluate` step's inference normaliser. Output: per-fold JSON files named `norm_stats_foldN.json` with mean and median across seeds for each feature dimension. Input file schema (one per fold/seed): { "fold": 0, "seed": 0, "feature_dim": 42, "mean": [, ...], # per-dim mean "std": [, ...], # per-dim std (sample stddev) "n_bars": 12345 # samples used to compute stats } Output file schema (one per fold): { "fold": N, "feature_dim": int, "n_seeds": int, "mean": [, ...], # mean of per-seed means "median": [, ...], # median of per-seed means "std": [, ...], # mean of per-seed stds "std_dispersion": [, ...], # cross-seed stddev of the means "n_bars_total": int } """ from __future__ import annotations import argparse import json import re import sys from collections import defaultdict from pathlib import Path from statistics import mean, median, pstdev NORM_STATS_RE = re.compile(r"norm_stats_fold(\d+)_seed(\d+)\.json") def main() -> int: ap = argparse.ArgumentParser( description=( "Merge norm_stats_foldN_seedM.json files across seeds into " "per-fold aggregates for inference normalisation." ), ) ap.add_argument( "--input-dir", type=Path, required=True, help="Directory containing norm_stats_foldN_seedM.json files", ) ap.add_argument( "--output-dir", type=Path, required=True, help="Directory for output norm_stats_foldN.json files", ) args = ap.parse_args() if not args.input_dir.is_dir(): print(f"ERROR: input dir not found: {args.input_dir}", file=sys.stderr) return 1 by_fold: dict[int, list[Path]] = defaultdict(list) for entry in args.input_dir.iterdir(): m = NORM_STATS_RE.match(entry.name) if m: by_fold[int(m.group(1))].append(entry) if not by_fold: print( f"ERROR: no norm_stats_foldN_seedM.json files found in {args.input_dir}", file=sys.stderr, ) return 1 args.output_dir.mkdir(parents=True, exist_ok=True) for fold, paths in sorted(by_fold.items()): loaded = [] for p in sorted(paths): try: with p.open("r", encoding="utf-8") as fh: loaded.append(json.load(fh)) except (OSError, json.JSONDecodeError) as e: print(f"WARN: skipping {p}: {e}", file=sys.stderr) if not loaded: print(f"WARN: no readable seeds for fold {fold}", file=sys.stderr) continue # Validate consistent feature_dim — abort if not (norm-stats files # for the same fold must agree on dimensionality). feature_dims = {entry.get("feature_dim") for entry in loaded} if len(feature_dims) != 1 or None in feature_dims: print( f"ERROR: fold {fold} has inconsistent feature_dim: {feature_dims}", file=sys.stderr, ) return 1 feature_dim = feature_dims.pop() # Per-dim aggregation: collect across seeds, then reduce. means_by_dim: list[list[float]] = [[] for _ in range(feature_dim)] stds_by_dim: list[list[float]] = [[] for _ in range(feature_dim)] for entry in loaded: entry_mean = entry["mean"] entry_std = entry["std"] if len(entry_mean) != feature_dim or len(entry_std) != feature_dim: print( f"ERROR: fold {fold} seed entry has mismatched dim", file=sys.stderr, ) return 1 for d in range(feature_dim): means_by_dim[d].append(float(entry_mean[d])) stds_by_dim[d].append(float(entry_std[d])) out = { "fold": fold, "feature_dim": feature_dim, "n_seeds": len(loaded), "mean": [mean(vals) for vals in means_by_dim], "median": [median(vals) for vals in means_by_dim], "std": [mean(vals) for vals in stds_by_dim], "std_dispersion": [ pstdev(vals) if len(vals) > 1 else 0.0 for vals in means_by_dim ], "n_bars_total": sum(int(entry.get("n_bars", 0)) for entry in loaded), } out_path = args.output_dir / f"norm_stats_fold{fold}.json" with out_path.open("w", encoding="utf-8") as fh: json.dump(out, fh, indent=2) print( f"fold {fold}: merged {len(loaded)} seeds → {out_path}", file=sys.stderr, ) return 0 if __name__ == "__main__": sys.exit(main())