""" CAISc Paper — 9-Model Mechanistic Decomposition Training on Modal. Trains 3 architectures × 3 ACh model variants: - v14a (synaptic suppression only) — existing data in prod_5k_v13 - v14b (depolarization only) — new data in v14b_5k - v14c (both mechanisms) — new data in v14c_5k Each variant trains Model B (with ACh level as input feature). Plus one Model A (no ACh) on v14a as the control baseline. Total: 9 Model B variants + 3 Model A controls = 12 models (but Model A is architecture-independent across data variants, so really 9 + 3 = 12). Actually: 3 arch × 3 data variants = 9 (Model B only, since all data has ACh sweeps) + 3 arch × 1 = 3 Model A (baseline, from v14a ACh=0 only) = 12 total training runs Usage: modal run train_decomposition.py # Train all 12 models modal run train_decomposition.py --experiment # Train + money experiment """ import modal app = modal.App("caisc-decomp-train") vol = modal.Volume.from_name("caisc-data") train_image = ( modal.Image.debian_slim(python_version="3.11") .pip_install( "torch>=2.2", "numpy", "scipy", "matplotlib", "scikit-learn", "xgboost", ) .add_local_dir("training", remote_path="/root/training") ) # Data paths for each ACh model variant DATA_DIRS = { "v14a": "/data/sims/prod_5k_v13/circuits", # synaptic suppression only "v14b": "/data/sims/v14b_5k/circuits", # depolarization only "v14c": "/data/sims/v14c_5k/circuits", # both mechanisms } @app.function( image=train_image, gpu="L4", volumes={"/data": vol}, timeout=3600, memory=16384, ) def train_single(data_variant: str, model_variant: str, arch: str) -> dict: """Train one model on one data variant with one architecture. Args: data_variant: "v14a", "v14b", or "v14c" — which ACh mechanism data model_variant: "A" (no ACh feature) or "B" (with ACh feature) arch: "mlp", "transformer", or "xgboost" """ import logging, sys, os, json logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s", handlers=[logging.StreamHandler(sys.stdout)]) logger = logging.getLogger(f"{data_variant}_{model_variant}_{arch}") sys.path.insert(0, "/root") from training.config import TrainConfig from training.train import train_one_model import torch device = torch.device("cuda" if torch.cuda.is_available() else "cpu") logger.info(f"Training {data_variant}/{model_variant}/{arch} on {device}") # Override config to point to correct data directory cfg = TrainConfig() cfg.sim_dir = DATA_DIRS[data_variant] # Use variant-specific checkpoint/log directories to avoid collisions cfg.checkpoint_dir = f"/data/training/decomp/checkpoints/{data_variant}" cfg.log_dir = f"/data/training/decomp/logs/{data_variant}" # Disable extra_ach0_dir (not relevant for decomposition) cfg.extra_ach0_dir = "" results = train_one_model(cfg, model_variant, device, arch=arch) # Tag results with data variant results["data_variant"] = data_variant results["arch"] = arch # Save combined results os.makedirs(f"/data/training/decomp/results", exist_ok=True) result_file = f"/data/training/decomp/results/{data_variant}_{model_variant}_{arch}.json" with open(result_file, "w") as f: json.dump(results, f, indent=2, default=str) vol.commit() logger.info(f"Done: {data_variant}/{model_variant}/{arch} — R²={results.get('mean_val_r2', '?')}") return results @app.function( image=train_image, gpu="L4", volumes={"/data": vol}, timeout=7200, memory=16384, ) def run_decomposition_experiment() -> dict: """Run the full mechanistic decomposition money experiment. For each of the 9 Model B variants (3 data × 3 arch), predict Allen V1 state changes and compare against real data. """ import logging, sys, os, json, glob import numpy as np logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s", handlers=[logging.StreamHandler(sys.stdout)]) logger = logging.getLogger("decomp_experiment") sys.path.insert(0, "/root") from training.config import TrainConfig, INPUT_FEATURES_B, OUTPUT_STATS from training.dataset import Normalizer from training.model import build_model_b import torch device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # Load Allen data allen_dir = "/data/allen/epochs" allen_files = sorted(glob.glob(f"{allen_dir}/allen_*.json")) logger.info(f"Allen files: {len(allen_files)}") # Group by session sessions = {} for fpath in allen_files: with open(fpath) as f: epoch = json.load(f) sid = epoch["session_id"] state = epoch["epoch_type"] # "running" or "stationary" if sid not in sessions: sessions[sid] = {"running": [], "stationary": []} if state in sessions[sid]: sessions[sid][state].append(epoch) # Filter to sessions with both states dual_sessions = {sid: data for sid, data in sessions.items() if data["running"] and data["stationary"]} logger.info(f"Sessions with both states: {len(dual_sessions)}") # Canonical circuit parameters for prediction canonical = { "n_exc": 160, "n_inh": 40, "conn_prob": 0.06, "n_synapses": 2400, "mean_in_degree": 12.0, "gS_exc_effective": 5e-6, "ou_mu_effective": -0.001, "ou_sigma_effective": 0.001, "ou_tau": 5.0, "sim_duration_ms": 3000.0, } all_results = {} # For each data variant × architecture, load Model B and run experiment for data_var in ["v14a", "v14b", "v14c"]: for arch in ["mlp", "transformer", "xgboost"]: key = f"{data_var}_B_{arch}" logger.info(f"\n{'='*60}") logger.info(f"EXPERIMENT: {key}") logger.info(f"{'='*60}") # Load checkpoint if arch == "xgboost": # XGBoost doesn't save standard checkpoints — skip for now # or load from results JSON logger.info(f" Skipping XGBoost (no checkpoint-based inference)") continue ckpt_path = f"/data/training/decomp/checkpoints/{data_var}/model_b/best.pt" if not os.path.exists(ckpt_path): logger.warning(f" Checkpoint not found: {ckpt_path}") continue ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) # Rebuild model cfg = TrainConfig() model_config = ckpt.get("config", {}) actual_arch = model_config.get("arch", arch) model = build_model_b(cfg, arch=actual_arch) model.load_state_dict(ckpt["model_state_dict"]) model = model.to(device) model.eval() # Rebuild normalizers x_norm = Normalizer.from_state_dict(ckpt["x_norm"]) y_norm = Normalizer.from_state_dict(ckpt["y_norm"]) # For each session, predict low vs high ACh session_results = [] for sid, data in dual_sessions.items(): # Mean observed stats running_stats = {} stationary_stats = {} for stat in OUTPUT_STATS: running_vals = [e["statistics"][stat] for e in data["running"] if stat in e.get("statistics", {})] stationary_vals = [e["statistics"][stat] for e in data["stationary"] if stat in e.get("statistics", {})] if running_vals and stationary_vals: running_stats[stat] = float(np.mean(running_vals)) stationary_stats[stat] = float(np.mean(stationary_vals)) if not running_stats: continue # Predict at low ACh (0.1) and high ACh (0.9) def make_input(ach_level): x = np.array([[canonical[f] for f in INPUT_FEATURES_B[:-1]] + [ach_level]]) return x_norm.transform(x) with torch.no_grad(): x_low = torch.tensor(make_input(0.1), dtype=torch.float32).to(device) x_high = torch.tensor(make_input(0.9), dtype=torch.float32).to(device) pred_low = model(x_low).cpu().numpy() pred_high = model(x_high).cpu().numpy() # Inverse normalize predictions pred_low_raw = y_norm.inverse(pred_low)[0] pred_high_raw = y_norm.inverse(pred_high)[0] # Compare directions sr = {} for i, stat in enumerate(OUTPUT_STATS): if stat in running_stats: obs_delta = running_stats[stat] - stationary_stats[stat] pred_delta = pred_high_raw[i] - pred_low_raw[i] sign_match = (np.sign(obs_delta) == np.sign(pred_delta)) sr[stat] = { "obs_delta": round(float(obs_delta), 6), "pred_delta": round(float(pred_delta), 6), "sign_match": bool(sign_match), } session_results.append({ "session_id": int(sid), "n_running": len(data["running"]), "n_stationary": len(data["stationary"]), "stats": sr, }) # Aggregate sign accuracy per stat sign_accuracy = {} for stat in OUTPUT_STATS: matches = [s["stats"][stat]["sign_match"] for s in session_results if stat in s["stats"]] if matches: sign_accuracy[stat] = round(sum(matches) / len(matches) * 100, 1) logger.info(f"\n Sign accuracy ({len(session_results)} sessions):") for stat, acc in sign_accuracy.items(): marker = "✅" if acc > 60 else "❌" if acc < 40 else " " logger.info(f" {stat:25s}: {acc:5.1f}% {marker}") mean_acc = float(np.mean(list(sign_accuracy.values()))) if sign_accuracy else 0 logger.info(f" {'MEAN':25s}: {mean_acc:5.1f}%") all_results[key] = { "data_variant": data_var, "arch": arch, "n_sessions": len(session_results), "sign_accuracy": sign_accuracy, "mean_sign_accuracy": round(mean_acc, 1), "session_details": session_results, } # Save full results os.makedirs("/data/training/decomp/results", exist_ok=True) with open("/data/training/decomp/results/decomposition_experiment.json", "w") as f: json.dump(all_results, f, indent=2, default=str) vol.commit() # Print summary table logger.info(f"\n\n{'='*80}") logger.info("MECHANISTIC DECOMPOSITION — SUMMARY TABLE") logger.info(f"{'='*80}") logger.info(f"{'Statistic':25s} {'v14a(syn)':>10s} {'v14b(dep)':>10s} {'v14c(both)':>10s}") logger.info(f"{'-'*25} {'-'*10} {'-'*10} {'-'*10}") for stat in OUTPUT_STATS: vals = [] for dv in ["v14a", "v14b", "v14c"]: # Use MLP results (or first available) for arch in ["mlp", "transformer"]: k = f"{dv}_B_{arch}" if k in all_results: vals.append(all_results[k]["sign_accuracy"].get(stat, 0)) break else: vals.append(0) logger.info(f"{stat:25s} {vals[0]:9.1f}% {vals[1]:9.1f}% {vals[2]:9.1f}%") logger.info(f"{'='*80}") return all_results @app.function( image=train_image, gpu="L4", volumes={"/data": vol}, timeout=10800, # 3 hours memory=16384, ) def orchestrate(run_experiment: bool = False) -> dict: """Launch all training jobs in parallel, then optionally run experiment.""" import logging, sys, time, json, os, glob logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s", handlers=[logging.StreamHandler(sys.stdout)]) logger = logging.getLogger("orchestrator") # Check data availability for name, path in DATA_DIRS.items(): import glob as g # Can't check from orchestrator (different container) — just log logger.info(f" {name}: {path}") # Launch all training jobs in parallel logger.info("=" * 70) logger.info("LAUNCHING PARALLEL TRAINING: 3 data × 3 arch × Model B + 3 Model A") logger.info("=" * 70) t0 = time.time() handles = {} # Model B: 3 data variants × 2 neural architectures (MLP, Transformer) # XGBoost separately (no GPU needed but uses same function) for data_var in ["v14a", "v14b", "v14c"]: for arch in ["mlp", "transformer", "xgboost"]: key = f"{data_var}_B_{arch}" handles[key] = train_single.spawn(data_var, "B", arch) logger.info(f" Spawned {key}") # Model A baseline: only from v14a (ACh=0 data), 3 architectures for arch in ["mlp", "transformer", "xgboost"]: key = f"v14a_A_{arch}" handles[key] = train_single.spawn("v14a", "A", arch) logger.info(f" Spawned {key}") logger.info(f"\n{len(handles)} jobs launched. Waiting for completion...\n") # Collect results results = {} for name, handle in handles.items(): logger.info(f"Waiting for {name}...") try: r = handle.get() results[name] = r logger.info(f" {name}: R²={r.get('mean_val_r2', '?')}, Time={r.get('total_time_s', 0):.0f}s") except Exception as e: logger.error(f" {name} FAILED: {e}") results[name] = {"error": str(e)} total_time = time.time() - t0 logger.info(f"\nAll training complete in {total_time:.0f}s ({total_time/60:.1f} min)") # Print summary table logger.info(f"\n{'='*90}") logger.info("MODEL B — Validation R² by Data Variant × Architecture") logger.info(f"{'='*90}") logger.info(f"{'Data':8s} {'Arch':15s} {'Mean R²':>10s} {'Best Epoch':>12s} {'Time(s)':>10s}") logger.info(f"{'-'*8} {'-'*15} {'-'*10} {'-'*12} {'-'*10}") for data_var in ["v14a", "v14b", "v14c"]: for arch in ["mlp", "transformer", "xgboost"]: key = f"{data_var}_B_{arch}" r = results.get(key, {}) if "error" in r: logger.info(f"{data_var:8s} {arch:15s} {'FAILED':>10s}") else: logger.info(f"{data_var:8s} {arch:15s} {r.get('mean_val_r2', 0):10.4f} {r.get('best_epoch', 0):12d} {r.get('total_time_s', 0):10.0f}") logger.info(f"\nModel A baselines:") for arch in ["mlp", "transformer", "xgboost"]: key = f"v14a_A_{arch}" r = results.get(key, {}) if "error" not in r: logger.info(f" A_{arch}: R²={r.get('mean_val_r2', 0):.4f}") logger.info(f"{'='*90}") vol.commit() # Save all results os.makedirs("/data/training/decomp/results", exist_ok=True) with open("/data/training/decomp/results/all_training_results.json", "w") as f: json.dump(results, f, indent=2, default=str) vol.commit() # Run experiment if requested if run_experiment: logger.info("\n\nRunning mechanistic decomposition experiment...") vol.reload() experiment_results = run_decomposition_experiment.remote() results["experiment"] = experiment_results with open("/data/training/decomp/results/all_results_with_experiment.json", "w") as f: json.dump(results, f, indent=2, default=str) vol.commit() return results @app.local_entrypoint() def main(experiment: bool = False, experiment_only: bool = False): """Launch the full mechanistic decomposition training + experiment. Args: experiment: Also run the money experiment after training. experiment_only: Skip training, just run the money experiment (requires prior training). """ if experiment_only: print("=" * 60) print("CAISc — EXPERIMENT ONLY (skipping training)") print("=" * 60) results = run_decomposition_experiment.remote() print(f"\nEXPERIMENT RESULTS:") for key, r in results.items(): if isinstance(r, dict) and "mean_sign_accuracy" in r: print(f" {key:25s}: {r['mean_sign_accuracy']:.1f}% mean sign accuracy") print(f"\nResults saved to Modal volume: /data/training/decomp/results/") return print("=" * 60) print("CAISc MECHANISTIC DECOMPOSITION") print(" Training: 3 data variants × 3 architectures × Model B") print(" + 3 Model A baselines") print(f" Experiment: {experiment}") print("=" * 60) results = orchestrate.remote(run_experiment=experiment) print(f"\n{'='*60}") print("TRAINING COMPLETE") print(f"{'='*60}") for key, r in results.items(): if key == "experiment": continue if isinstance(r, dict) and "error" not in r: print(f" {key:25s}: R²={r.get('mean_val_r2', '?'):>8s}, Epoch={r.get('best_epoch', '?')}") elif isinstance(r, dict): print(f" {key:25s}: FAILED — {r.get('error', '?')[:50]}") if "experiment" in results: print(f"\nEXPERIMENT RESULTS:") exp = results["experiment"] for key, r in exp.items(): if isinstance(r, dict) and "mean_sign_accuracy" in r: print(f" {key:25s}: {r['mean_sign_accuracy']:.1f}% mean sign accuracy") print(f"\nResults saved to Modal volume: /data/training/decomp/results/")