| """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) |
|
|
|
|
| |
| |
| |
|
|
| 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") |
|
|
| |
| 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", "") |
|
|
| |
| if not preprocessed_path and h5ad_path: |
| candidate = h5ad_path.replace(".h5ad", "_preprocessed.pt") |
| if os.path.exists(candidate): |
| preprocessed_path = candidate |
|
|
| |
| |
| 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}") |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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" |
|
|
| |
| 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}") |
|
|
| |
| 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) |
|
|
| |
| 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: |
| |
| 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"] |
| dose = batch["dose"].to(device) |
|
|
| |
| known_targets = perturbation |
|
|
| |
| |
| 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 = 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: |
| |
| |
| 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: |
| |
| return {"target": 0.0, "contrastive": 0.0, "dose": 0.0} |
|
|
| |
| 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) |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
| |
| 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_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), |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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}") |
|
|
| |
| 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"]) |
|
|
| |
| print("\n=== Building model ===") |
|
|
| |
| 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}") |
|
|
| |
| 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), |
| ) |
|
|
| |
| 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}") |
|
|
| |
| 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") |
|
|
| |
| 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"]): |
| |
| 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 |
|
|
| |
| train_crispr = True |
| 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)") |
|
|
| |
| 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}") |
|
|
| |
| 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() |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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}") |
|
|
| |
| if epochs_no_improve >= patience: |
| print(f"\nEarly stopping: no R@1 improvement for {patience} epochs. Best={best_train_r1:.3f}") |
| break |
|
|
| |
| if model.use_drug_gene_bridge: |
| gate_val = torch.sigmoid(model.drug_gate).item() |
| print(f" Drug gate: {gate_val:.4f}") |
|
|
| |
| 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() |
|
|