| """ |
| 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_DIRS = { |
| "v14a": "/data/sims/prod_5k_v13/circuits", |
| "v14b": "/data/sims/v14b_5k/circuits", |
| "v14c": "/data/sims/v14c_5k/circuits", |
| } |
|
|
|
|
| @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}") |
| |
| |
| cfg = TrainConfig() |
| cfg.sim_dir = DATA_DIRS[data_variant] |
| |
| |
| cfg.checkpoint_dir = f"/data/training/decomp/checkpoints/{data_variant}" |
| cfg.log_dir = f"/data/training/decomp/logs/{data_variant}" |
| |
| |
| cfg.extra_ach0_dir = "" |
| |
| results = train_one_model(cfg, model_variant, device, arch=arch) |
| |
| |
| results["data_variant"] = data_variant |
| results["arch"] = arch |
| |
| |
| 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") |
| |
| |
| allen_dir = "/data/allen/epochs" |
| allen_files = sorted(glob.glob(f"{allen_dir}/allen_*.json")) |
| logger.info(f"Allen files: {len(allen_files)}") |
| |
| |
| sessions = {} |
| for fpath in allen_files: |
| with open(fpath) as f: |
| epoch = json.load(f) |
| sid = epoch["session_id"] |
| state = epoch["epoch_type"] |
| if sid not in sessions: |
| sessions[sid] = {"running": [], "stationary": []} |
| if state in sessions[sid]: |
| sessions[sid][state].append(epoch) |
| |
| |
| 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 = { |
| "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 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}") |
| |
| |
| if arch == "xgboost": |
| |
| |
| 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) |
| |
| |
| 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() |
| |
| |
| x_norm = Normalizer.from_state_dict(ckpt["x_norm"]) |
| y_norm = Normalizer.from_state_dict(ckpt["y_norm"]) |
| |
| |
| session_results = [] |
| for sid, data in dual_sessions.items(): |
| |
| 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 |
| |
| |
| 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() |
| |
| |
| pred_low_raw = y_norm.inverse(pred_low)[0] |
| pred_high_raw = y_norm.inverse(pred_high)[0] |
| |
| |
| 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, |
| }) |
| |
| |
| 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, |
| } |
| |
| |
| 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() |
| |
| |
| 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"]: |
| |
| 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, |
| 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") |
| |
| |
| for name, path in DATA_DIRS.items(): |
| import glob as g |
| |
| logger.info(f" {name}: {path}") |
| |
| |
| logger.info("=" * 70) |
| logger.info("LAUNCHING PARALLEL TRAINING: 3 data × 3 arch × Model B + 3 Model A") |
| logger.info("=" * 70) |
| |
| t0 = time.time() |
| handles = {} |
| |
| |
| |
| 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}") |
| |
| |
| 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") |
| |
| |
| 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)") |
| |
| |
| 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() |
| |
| |
| 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() |
| |
| |
| 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/") |
|
|