GID-Flow / PDGrapher /scripts /train_causal_flow_drug.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
29.2 kB
"""Joint CRISPRi + Drug training for CausalFlow-ID.
Two-stage training schedule:
Stage 1 (CRISPRi pretraining):
Epochs 0 β†’ drug_start_epoch-1
Only CRISPRi data, no drug losses
Learns basic population encoding, flow matching, causal gene identification
Stage 2 (Drug warmup):
Epochs drug_start_epoch β†’ drug_start_epoch + drug_warmup_epochs
Both CRISPRi + SciPlex3 data
Drug loss weights linearly ramped from 0 β†’ target values
Stage 3 (Joint training):
Epochs after warmup
Both datasets, full drug loss weights
Curriculum learning (Phase 1-4):
Epochs 1-10: Flow matching only
Epochs 11-20: Add cycle loss
Epochs 21-30: Add causal + sparse loss
Epochs 31+: Full objective
Usage:
# CRISPRi pretraining only (Stage 1)
python scripts/train_causal_flow_drug.py --config configs/causal_flow_drug.yaml --stage crispr_only
# Full two-stage training
python scripts/train_causal_flow_drug.py --config configs/causal_flow_drug.yaml
# Debug: 3-epoch smoke test
python scripts/train_causal_flow_drug.py --config configs/causal_flow_drug.yaml --debug
"""
import argparse
import csv
import os
import sys
import warnings
from typing import Any, Dict, Optional
warnings.filterwarnings("ignore")
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))
import torch
import yaml
from gidflow.data import (
ScPerturbPopulationDataset,
Sciplex3Dataset,
population_collate_fn,
sciplex_collate_fn,
)
from gidflow.losses.causal_flow_loss import CurriculumScheduler
from gidflow.metrics import compute_all_target_metrics
from gidflow.models import CausalFlowGIDModel
from gidflow.utils import set_seed, save_checkpoint, load_checkpoint
def load_config(path: str) -> dict:
with open(path) as f:
return yaml.safe_load(f)
# ──────────────────────────────────────────────────────────────────────────────
# Dataset builders
# ──────────────────────────────────────────────────────────────────────────────
def build_crispr_dataset(cfg: dict, seed: int):
"""Build CRISPRi (Norman2019) dataset and loaders."""
ds_cfg = cfg["dataset"]
ds = ScPerturbPopulationDataset(
h5ad_path=ds_cfg["h5ad_path"],
n_hvg=ds_cfg.get("n_hvg", 2000),
min_cells_per_cond=ds_cfg.get("min_cells_per_cond", 30),
max_source_cells=ds_cfg.get("max_source_cells", 64),
max_target_cells=ds_cfg.get("max_target_cells", 64),
control_label=ds_cfg.get("control_label", "control"),
use_single_pert_only=not ds_cfg.get("include_combos_in_train", False),
force_include_pert_genes=ds_cfg.get("force_include_pert_genes", True),
include_combos_in_train=ds_cfg.get("include_combos_in_train", False),
target_sum=ds_cfg.get("target_sum", 1e4),
seed=seed,
)
num_genes = ds.num_genes
print(f" CRISPRi dataset: {len(ds)} conditions, {num_genes} genes")
# Gene-disjoint split
train_idx, val_idx = ds.get_gene_disjoint_split(
val_fraction=cfg["train"].get("val_split", 0.15),
seed=seed,
)
from torch.utils.data import Subset
train_ds = Subset(ds, train_idx)
val_ds = Subset(ds, val_idx)
print(f" Gene-disjoint split: train={len(train_idx)}, val={len(val_idx)}")
bs = cfg["train"]["batch_size"]
train_loader = torch.utils.data.DataLoader(
train_ds, batch_size=bs, shuffle=True,
collate_fn=population_collate_fn, num_workers=0,
)
val_loader = torch.utils.data.DataLoader(
val_ds, batch_size=bs, shuffle=False,
collate_fn=population_collate_fn, num_workers=0,
)
return ds, train_loader, val_loader, num_genes
def build_drug_dataset(cfg: dict, num_genes: int, seed: int):
"""Build SciPlex3 drug dataset and loaders.
Uses preprocessed .pt file if available for speed.
"""
ds_cfg = cfg["dataset"]
sc_cfg = ds_cfg.get("sciplex", {})
h5ad_path = sc_cfg.get("h5ad_path", "")
preprocessed_path = sc_cfg.get("preprocessed_path", "")
# Auto-generate preprocessed path from h5ad_path if not specified
if not preprocessed_path and h5ad_path:
candidate = h5ad_path.replace(".h5ad", "_preprocessed.pt")
if os.path.exists(candidate):
preprocessed_path = candidate
# If no preprocessed file exists, try to create one on-the-fly
# (this requires enough RAM)
if not preprocessed_path:
preprocessed_path = None
print(f" SciPlex3 h5ad: {h5ad_path}")
print(f" SciPlex3 preprocessed: {preprocessed_path}")
ds = Sciplex3Dataset(
h5ad_path=h5ad_path if not preprocessed_path else None,
n_hvg=ds_cfg.get("n_hvg", 2000),
min_cells_per_cond=sc_cfg.get("min_cells_per_cond", 30),
max_source_cells=sc_cfg.get("max_source_cells", 64),
max_target_cells=sc_cfg.get("max_target_cells", 64),
cell_lines=sc_cfg.get("cell_lines"),
doses=sc_cfg.get("doses"),
times=sc_cfg.get("times"),
target_sum=sc_cfg.get("target_sum", 1e4),
seed=seed,
drug_emb_dim=sc_cfg.get("drug_emb_dim", 128),
preprocessed_path=preprocessed_path,
target_num_genes=num_genes,
drug_smiles_csv=sc_cfg.get("drug_smiles_csv", ""),
)
print(f" SciPlex3 dataset: {len(ds)} conditions, {ds.num_genes} genes")
print(f" Unique drugs: {len(ds.unique_drugs)}")
print(f" Cell lines: {ds.cell_line_list}")
# Random split for drug data
n_val = max(1, int(len(ds) * cfg["train"].get("val_split", 0.15)))
n_train = len(ds) - n_val
from torch.utils.data import random_split
train_ds, val_ds = random_split(
ds, [n_train, n_val],
generator=torch.Generator().manual_seed(seed),
)
print(f" Drug split: train={n_train}, val={n_val}")
bs = cfg["train"]["batch_size"]
train_loader = torch.utils.data.DataLoader(
train_ds, batch_size=bs, shuffle=True,
collate_fn=lambda b: sciplex_collate_fn(b, num_genes=num_genes), num_workers=0,
)
val_loader = torch.utils.data.DataLoader(
val_ds, batch_size=bs, shuffle=False,
collate_fn=lambda b: sciplex_collate_fn(b, num_genes=num_genes), num_workers=0,
)
return ds, train_loader, val_loader
# ──────────────────────────────────────────────────────────────────────────────
# Training loop
# ──────────────────────────────────────────────────────────────────────────────
def run_crispr_epoch(model, loader, optimizer, device, train: bool, epoch: int = 0) -> dict:
"""Run one epoch on CRISPRi data (no drug conditioning)."""
model.train(train)
totals = {"total": 0.0, "flow": 0.0, "cycle": 0.0,
"causal": 0.0, "sparse": 0.0, "recon": 0.0}
tgt_scores_all, tgt_true_all = [], []
n_batches = 0
for batch in loader:
batch = batch.to(device)
true_pert = batch.perturbation.float()
out = model(
batch.source_cells, batch.target_cells,
batch.source_mask, batch.target_mask,
true_perturbation=true_pert,
epoch=epoch,
)
total_loss = out["total_loss"]
if train:
optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
for k in totals:
totals[k] += out[f"{k}_loss"].item()
n_batches += 1
tgt_scores_all.append(out["target_scores"].detach())
tgt_true_all.append(true_pert.detach())
avg = {k: v / max(n_batches, 1) for k, v in totals.items()}
avg["mode"] = "crispr"
# Perturbation recovery metrics
pert_metrics = {}
if tgt_scores_all:
all_scores = torch.cat(tgt_scores_all, dim=0)
all_targets = torch.cat(tgt_true_all, dim=0)
try:
tgt_m = compute_all_target_metrics(all_scores, all_targets, ks=(1, 5, 10))
pert_metrics = {
"recall@1": tgt_m.get("recall@1", 0.0),
"recall@5": tgt_m.get("recall@5", 0.0),
"recall@10": tgt_m.get("recall@10", 0.0),
"precision@1": tgt_m.get("precision@1", 0.0),
"ndcg@1": tgt_m.get("ndcg@1", 0.0),
"mrr": tgt_m.get("mrr", 0.0),
}
except Exception as e:
print(f" Warning: CRISPRi metrics failed: {e}")
# Causal score diagnostics
avg["causal_mean"] = all_scores.mean().item()
avg["causal_std"] = all_scores.std().item()
avg["causal_max"] = all_scores.max().item()
avg["causal_min"] = all_scores.min().item()
return {**avg, **pert_metrics}
def run_drug_epoch(model, loader, optimizer, device, train: bool,
drug_loss_weights: Optional[Dict[str, float]] = None,
epoch: int = 0) -> dict:
"""Run one epoch on SciPlex3 drug data.
Parameters
----------
drug_loss_weights : dict with keys target, contrastive, dose
Weights for drug alignment loss components. If None, use full weights.
epoch : current epoch for curriculum scheduling.
"""
model.train(train)
# Default totals
totals = {"total": 0.0, "flow": 0.0, "cycle": 0.0,
"causal": 0.0, "sparse": 0.0, "recon": 0.0,
"drug_total": 0.0, "drug_target": 0.0,
"drug_contrastive": 0.0, "drug_contrastive_diversity": 0.0,
"drug_dose": 0.0}
n_batches = 0
for batch in loader:
# batch is a dict from sciplex_collate_fn
source_cells = batch["source_cells"].to(device)
target_cells = batch["target_cells"].to(device)
source_mask = batch["source_mask"].to(device)
target_mask = batch["target_mask"].to(device)
perturbation = batch["perturbation"].to(device)
drug_smiles = batch["drug_smiles"] # list of str
dose = batch["dose"].to(device)
# Build known drug targets (use perturbation vector from dataset)
known_targets = perturbation
# Extract drug names for dose consistency
# sciplex_collate_fn returns metadata as list of dicts
meta_list = batch.get("metadata", [])
if meta_list and isinstance(meta_list, list):
drug_names = [b.get("drug_name", "") for b in meta_list]
else:
drug_names = None
out = model(
source_cells, target_cells,
source_mask, target_mask,
true_perturbation=perturbation,
drug_smiles=drug_smiles,
dose=dose,
known_drug_targets=known_targets,
drug_name=drug_names,
epoch=epoch,
)
# total_loss already includes drug losses (weighted by model's lambda)
total_loss = out["total_loss"]
if train:
optimizer.zero_grad()
total_loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
for k in totals:
# Map totals key β†’ model output key
# "total" β†’ "total_loss", "drug_total" β†’ "drug_total_loss", etc.
out_key = f"{k}_loss" if k in totals else k
if out_key in out:
totals[k] += out[out_key].item()
elif k in out:
totals[k] += out[k].item()
n_batches += 1
avg = {k: v / max(n_batches, 1) for k, v in totals.items()}
avg["mode"] = "drug"
return avg
def compute_drug_loss_weights(cfg: dict, epoch: int) -> Dict[str, float]:
"""Compute drug loss weights based on training stage.
Stage 1 (epoch < drug_start_epoch): all weights = 0
Stage 2 (drug_start_epoch ≀ epoch < drug_start_epoch + warmup): linear ramp
Stage 3 (epoch β‰₯ drug_start_epoch + warmup): full weights
"""
t_cfg = cfg["train"]
m_cfg = cfg["model"]
start_epoch = t_cfg.get("drug_start_epoch", 50)
warmup_epochs = t_cfg.get("drug_warmup_epochs", 20)
if epoch < start_epoch:
# CRISPRi only β€” no drug losses
return {"target": 0.0, "contrastive": 0.0, "dose": 0.0}
# Progress through warmup
progress = min(1.0, (epoch - start_epoch) / max(warmup_epochs, 1))
return {
"target": m_cfg.get("lambda_drug_target", 1.0) * progress,
"contrastive": m_cfg.get("lambda_drug_contrastive", 0.1) * progress,
"dose": m_cfg.get("lambda_drug_dose", 0.01) * progress,
}
def log_metrics(writer, epoch: int, split: str, metrics: dict):
"""Write metrics row to CSV."""
row = {"epoch": epoch, "mode": split}
for k, v in metrics.items():
if k in ("mode",):
continue
if isinstance(v, float):
row[k] = f"{v:.4f}"
else:
row[k] = str(v)
writer.writerow(row)
# ──────────────────────────────────────────────────────────────────────────────
# Model builder
# ──────────────────────────────────────────────────────────────────────────────
def build_model(cfg: dict, num_genes: int, device: torch.device,
curriculum_scheduler: Optional[CurriculumScheduler] = None) -> CausalFlowGIDModel:
"""Build CausalFlowGIDModel from config."""
m_cfg = cfg["model"]
m_cfg = dict(m_cfg)
m_cfg["num_genes"] = num_genes
# Drug config
sc_cfg = cfg["dataset"].get("sciplex", {})
use_drug = sc_cfg.get("enabled", False)
model = CausalFlowGIDModel(
num_genes=num_genes,
encoder_hidden=m_cfg["encoder_hidden"],
encoder_output=m_cfg["encoder_output"],
gap_hidden=m_cfg["gap_hidden"],
gap_output=m_cfg["gap_output"],
causal_gene_emb_dim=m_cfg["causal_gene_emb_dim"],
causal_n_heads=m_cfg["causal_n_heads"],
causal_n_layers=m_cfg["causal_n_layers"],
planner_hidden=m_cfg["planner_hidden"],
planner_n_layers=m_cfg.get("planner_n_layers", 2),
planner_topk=m_cfg.get("planner_topk"),
flow_latent_dim=m_cfg["flow_latent_dim"],
flow_hidden_dim=m_cfg["flow_hidden_dim"],
flow_n_layers=m_cfg["flow_n_layers"],
flow_time_embed_dim=m_cfg.get("flow_time_embed_dim", 128),
flow_pert_emb_dim=m_cfg.get("flow_pert_emb_dim", 256),
n_layers=m_cfg.get("n_layers", 2),
use_cooccurrence=m_cfg.get("use_cooccurrence", True),
use_latent=m_cfg.get("use_latent", True),
# Drug
drug_emb_dim=sc_cfg.get("drug_emb_dim", 128),
use_drug_encoder=use_drug,
use_drug_gene_bridge=use_drug,
drug_gate_init=m_cfg.get("drug_gate_init", 0.1),
# Loss weights
lambda_flow=m_cfg.get("lambda_flow", 1.0),
lambda_cycle=m_cfg.get("lambda_cycle", 0.5),
lambda_causal=m_cfg.get("lambda_causal", 0.1),
lambda_sparse=m_cfg.get("lambda_sparse", 0.01),
lambda_recon=m_cfg.get("lambda_recon", 0.1),
lambda_drug_target=m_cfg.get("lambda_drug_target", 1.0),
lambda_drug_contrastive=m_cfg.get("lambda_drug_contrastive", 0.1),
lambda_drug_dose=m_cfg.get("lambda_drug_dose", 0.01),
sparse_variance_weight=m_cfg.get("sparse_variance_weight", 0.0),
use_causal_infonce=m_cfg.get("use_causal_infonce", False),
infonce_neg_samples=m_cfg.get("infonce_neg_samples", 64),
infonce_temperature=m_cfg.get("infonce_temperature", 0.1),
curriculum_scheduler=curriculum_scheduler,
).to(device)
return model
# ──────────────────────────────────────────────────────────────────────────────
# Main
# ──────────────────────────────────────────────────────────────────────────────
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True)
parser.add_argument("--resume", default=None)
parser.add_argument("--debug", action="store_true")
parser.add_argument("--stage", choices=["crispr_only", "drug_only", "joint"], default=None)
parser.add_argument("--skip-drug-data", action="store_true",
help="Skip loading SciPlex3 data (CRISPRi-only mode)")
args = parser.parse_args()
cfg = load_config(args.config)
if args.debug:
cfg["train"]["epochs"] = 3
cfg["train"]["drug_start_epoch"] = 0
cfg["train"]["drug_warmup_epochs"] = 2
print("DEBUG mode: 3 epochs, drug losses start immediately")
set_seed(cfg["seed"])
device = torch.device(cfg["train"].get("device", "cuda") if torch.cuda.is_available() else "cpu")
print(f"Device: {device}")
print(f"Config: {args.config}")
# ── Datasets ──────────────────────────────────────────────────────────
print("\n=== Loading CRISPRi dataset ===")
crispr_ds, crispr_train_loader, crispr_val_loader, num_genes = build_crispr_dataset(cfg, cfg["seed"])
drug_train_loader = None
drug_val_loader = None
sc_cfg = cfg["dataset"].get("sciplex", {})
use_drug = sc_cfg.get("enabled", False) and not args.skip_drug_data
if use_drug:
print("\n=== Loading SciPlex3 drug dataset ===")
drug_ds, drug_train_loader, drug_val_loader = build_drug_dataset(cfg, num_genes, cfg["seed"])
# ── Model ─────────────────────────────────────────────────────────────
print("\n=== Building model ===")
# Curriculum learning scheduler for dynamic loss weights
curriculum_scheduler = CurriculumScheduler(
flow_end_epoch=t_cfg.get("curriculum_flow_end_epoch", 10),
cycle_start_epoch=t_cfg.get("curriculum_cycle_start_epoch", 11),
cycle_end_epoch=t_cfg.get("curriculum_cycle_end_epoch", 20),
causal_start_epoch=t_cfg.get("curriculum_causal_start_epoch", 21),
causal_end_epoch=t_cfg.get("curriculum_causal_end_epoch", 30),
)
print(f" Curriculum schedule: flow-only→cycle→causal (epochs 1-10→11-20→21-30)")
model = build_model(cfg, num_genes, device, curriculum_scheduler=curriculum_scheduler)
total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f" Total trainable parameters: {total_params:,}")
print(f" Drug encoder: {'enabled' if model.use_drug_encoder else 'disabled'}")
print(f" Drug gene bridge: {'enabled' if model.use_drug_gene_bridge else 'disabled'}")
if model.use_drug_gene_bridge:
print(f" Drug gate initial value: {torch.sigmoid(model.drug_gate).item():.4f}")
# ── Optimizer ─────────────────────────────────────────────────────────
t_cfg = cfg["train"]
optimizer = torch.optim.AdamW(
model.parameters(),
lr=t_cfg["lr"],
weight_decay=t_cfg.get("weight_decay", 1e-4),
)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=t_cfg["epochs"], eta_min=t_cfg.get("lr_min", 1e-6),
)
# ── Resume ────────────────────────────────────────────────────────────
start_epoch = 0
if args.resume and os.path.exists(args.resume):
ckpt = load_checkpoint(args.resume, model, optimizer, device)
start_epoch = ckpt.get("epoch", 0) + 1
print(f"Resumed from epoch {start_epoch - 1}")
# ── Override stage ────────────────────────────────────────────────────
if args.stage == "crispr_only":
use_drug = False
print("Stage override: CRISPRi only (no drug data)")
elif args.stage == "drug_only":
if drug_train_loader is None:
print("ERROR: drug data not loaded. Remove --skip-drug-data")
return
print("Stage override: Drug only")
# ── Output setup ──────────────────────────────────────────────────────
out_dir = cfg["output"]["dir"]
ckpt_path = cfg["output"]["checkpoint"]
best_ckpt_path = cfg["output"].get("best_checkpoint", ckpt_path.replace(".pt", "_best.pt"))
metrics_csv = cfg["output"]["metrics_csv"]
os.makedirs(out_dir, exist_ok=True)
fieldnames = [
"epoch", "mode",
"total", "flow", "cycle", "causal", "sparse", "recon",
"recall@1", "recall@5", "recall@10",
"precision@1", "ndcg@1", "mrr",
"drug_total", "drug_target", "drug_contrastive", "drug_contrastive_diversity", "drug_dose",
"drug_w_target", "drug_w_contrastive", "drug_w_dose",
"causal_mean", "causal_std", "causal_max", "causal_min",
]
with open(metrics_csv, "w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
best_train_r1 = 0.0
patience = t_cfg.get("early_stopping_patience", 30)
epochs_no_improve = 0
start_epoch = max(start_epoch, 0)
for epoch in range(start_epoch, t_cfg["epochs"]):
# ── Compute drug loss weights for this epoch ──────────────────
drug_weights = compute_drug_loss_weights(cfg, epoch)
w_target = drug_weights["target"]
w_contrastive = drug_weights["contrastive"]
w_dose = drug_weights["dose"]
use_drug_this_epoch = w_target > 0 or w_contrastive > 0 or w_dose > 0
# ── Determine which datasets to train on ─────────────────────
train_crispr = True # Always train on CRISPRi
train_drug = use_drug and use_drug_this_epoch
print(f"\n--- Epoch {epoch} ---")
if train_drug:
print(f" Drug weights: target={w_target:.3f} contrastive={w_contrastive:.4f} dose={w_dose:.5f}")
else:
print(" Mode: CRISPRi pretraining (no drug losses)")
# ── Train CRISPRi ─────────────────────────────────────────────
if train_crispr:
crispr_train_m = run_crispr_epoch(
model, crispr_train_loader, optimizer, device, train=True,
epoch=epoch,
)
crispr_train_m["drug_w_target"] = w_target
crispr_train_m["drug_w_contrastive"] = w_contrastive
crispr_train_m["drug_w_dose"] = w_dose
log_metrics(writer, epoch, "crispr_train", crispr_train_m)
cs_mean = crispr_train_m.get("causal_mean", 0)
cs_std = crispr_train_m.get("causal_std", 0)
r1 = crispr_train_m.get("recall@1", 0.0)
print(f" CRISPRi train | R@1={r1:.3f} "
f"loss={crispr_train_m['total']:.4f} "
f"flow={crispr_train_m['flow']:.4f} "
f"causal={crispr_train_m['causal']:.4f} "
f"cs={cs_mean:.3f}Β±{cs_std:.3f}")
# ── Train Drug ────────────────────────────────────────────────
if train_drug and drug_train_loader is not None:
drug_train_m = run_drug_epoch(
model, drug_train_loader, optimizer, device, train=True,
drug_loss_weights={"target": w_target, "contrastive": w_contrastive, "dose": w_dose},
epoch=epoch,
)
drug_train_m["drug_w_target"] = w_target
drug_train_m["drug_w_contrastive"] = w_contrastive
drug_train_m["drug_w_dose"] = w_dose
log_metrics(writer, epoch, "drug_train", drug_train_m)
print(f" Drug train | total={drug_train_m['total']:.4f} "
f"flow={drug_train_m['flow']:.4f} "
f"drug_t={drug_train_m.get('drug_target', 0):.4f} "
f"drug_c={drug_train_m.get('drug_contrastive', 0):.4f} "
f"drug_cd={drug_train_m.get('drug_contrastive_diversity', 0):.4f} "
f"drug_d={drug_train_m.get('drug_dose', 0):.4f}")
scheduler.step()
# ── Validate ─────────────────────────────────────────────────
with torch.no_grad():
if crispr_val_loader is not None:
crispr_val_m = run_crispr_epoch(
model, crispr_val_loader, optimizer, device, train=False,
epoch=epoch,
)
crispr_val_m["drug_w_target"] = w_target
crispr_val_m["drug_w_contrastive"] = w_contrastive
crispr_val_m["drug_w_dose"] = w_dose
log_metrics(writer, epoch, "crispr_val", crispr_val_m)
val_r1 = crispr_val_m.get("recall@1", 0.0)
if train_drug and drug_val_loader is not None:
drug_val_m = run_drug_epoch(
model, drug_val_loader, optimizer, device, train=False,
drug_loss_weights={"target": w_target, "contrastive": w_contrastive, "dose": w_dose},
epoch=epoch,
)
drug_val_m["drug_w_target"] = w_target
drug_val_m["drug_w_contrastive"] = w_contrastive
drug_val_m["drug_w_dose"] = w_dose
log_metrics(writer, epoch, "drug_val", drug_val_m)
# ── Checkpoint & early stopping ───────────────────────────────
train_r1 = crispr_train_m.get("recall@1", 0.0) if train_crispr else 0.0
if train_r1 > best_train_r1:
best_train_r1 = train_r1
epochs_no_improve = 0
save_checkpoint(
best_ckpt_path, model, optimizer,
epoch=epoch,
metrics={"best_train_r1": best_train_r1, "epoch": epoch},
config=cfg,
)
else:
epochs_no_improve += 1
# Periodic checkpoint
if (epoch + 1) % t_cfg.get("save_every", 10) == 0:
save_checkpoint(
ckpt_path, model, optimizer,
epoch=epoch,
metrics={"best_train_r1": best_train_r1, "epoch": epoch},
config=cfg,
)
print(f" Periodic checkpoint: {ckpt_path}")
# Early stopping
if epochs_no_improve >= patience:
print(f"\nEarly stopping: no R@1 improvement for {patience} epochs. Best={best_train_r1:.3f}")
break
# Drug gate diagnostic
if model.use_drug_gene_bridge:
gate_val = torch.sigmoid(model.drug_gate).item()
print(f" Drug gate: {gate_val:.4f}")
# Final save
save_checkpoint(
ckpt_path, model, optimizer,
epoch=epoch,
metrics={"best_train_r1": best_train_r1, "epoch": epoch},
config=cfg,
)
print(f"\nTraining complete. Best train R@1={best_train_r1:.3f}")
print(f" Best checkpoint: {best_ckpt_path}")
print(f" Metrics: {metrics_csv}")
if __name__ == "__main__":
main()