"""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()