caisc-2026-ach-transfer / code /train_decomposition.py
ronniebasak's picture
Upload code/train_decomposition.py with huggingface_hub
8e42904 verified
Raw
History Blame Contribute Delete
18.4 kB
"""
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/")