#!/usr/bin/env python3 """Train and evaluate the MitoInteract v2 head over cached frozen embeddings.""" from __future__ import annotations import argparse import copy import json import math import random from pathlib import Path import numpy as np import torch from safetensors.torch import save_file from torch.utils.data import DataLoader, TensorDataset from mitointeract_recovery.metrics import regression_metrics from mitointeract_recovery.model import MitoInteractHead, TargetScaler def read_manifest(path: Path) -> dict[str, str]: with path.open() as handle: return { row["pair_id"]: row["split"] for row in (json.loads(line) for line in handle if line.strip()) } def evaluate( model: MitoInteractHead, protein: torch.Tensor, ligand: torch.Tensor, targets: torch.Tensor, indices: np.ndarray, scaler: TargetScaler, ) -> dict: model.eval() with torch.inference_mode(): prediction = ( scaler.decode(model(protein[indices], ligand[indices])).cpu().numpy() ) return regression_metrics(targets[indices].cpu().numpy(), prediction) def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--embeddings", type=Path, required=True) parser.add_argument("--manifest", type=Path, required=True) parser.add_argument("--output-dir", type=Path, default=Path("artifacts/v2-head")) parser.add_argument("--target-name") parser.add_argument("--epochs", type=int, default=100) parser.add_argument("--batch-size", type=int, default=256) parser.add_argument("--learning-rate", type=float, default=1e-3) parser.add_argument("--weight-decay", type=float, default=1e-4) parser.add_argument("--warmup-ratio", type=float, default=0.05) parser.add_argument("--patience", type=int, default=10) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--projection-dim", type=int, default=256) parser.add_argument("--hidden-dim", type=int, default=512) parser.add_argument("--dropout", type=float, default=0.1) parser.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto") args = parser.parse_args() random.seed(args.seed) np.random.seed(args.seed) torch.manual_seed(args.seed) device_name = ( "cuda" if args.device == "auto" and torch.cuda.is_available() else args.device ) if device_name == "auto": device_name = "cpu" device = torch.device(device_name) arrays = np.load(args.embeddings) pair_ids = arrays["pair_ids"].astype(str) protein = torch.from_numpy(arrays["protein"]).float().to(device) ligand = torch.from_numpy(arrays["ligand"]).float().to(device) targets = torch.from_numpy(arrays["target"]).float().to(device) manifest = read_manifest(args.manifest) split_indices = { split: np.asarray( [ index for index, pair_id in enumerate(pair_ids) if manifest.get(pair_id) == split ] ) for split in ("train", "validation", "test") } if any(not len(indices) for indices in split_indices.values()): raise ValueError( f"all splits must be non-empty: { {k: len(v) for k, v in split_indices.items()} }" ) scaler = TargetScaler.fit(targets[split_indices["train"]]) encoded_targets = scaler.encode(targets) train_dataset = TensorDataset( protein[split_indices["train"]], ligand[split_indices["train"]], encoded_targets[split_indices["train"]], ) generator = torch.Generator().manual_seed(args.seed) loader = DataLoader( train_dataset, batch_size=min(args.batch_size, len(train_dataset)), shuffle=True, generator=generator, ) model = MitoInteractHead( protein.shape[1], ligand.shape[1], projection_dim=args.projection_dim, hidden_dim=args.hidden_dim, dropout=args.dropout, ).to(device) optimizer = torch.optim.AdamW( model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay, ) total_steps = max(1, args.epochs * len(loader)) warmup_steps = max(1, round(total_steps * args.warmup_ratio)) def lr_factor(step: int) -> float: if step < warmup_steps: return (step + 1) / warmup_steps progress = (step - warmup_steps) / max(1, total_steps - warmup_steps) return 0.5 * (1 + math.cos(math.pi * progress)) scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_factor) best_state = None best_validation = math.inf best_epoch = 0 remaining_patience = args.patience history = [] for epoch in range(1, args.epochs + 1): model.train() losses = [] for batch_protein, batch_ligand, batch_target in loader: optimizer.zero_grad(set_to_none=True) prediction = model(batch_protein, batch_ligand) loss = torch.nn.functional.mse_loss(prediction, batch_target) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() losses.append(float(loss.detach())) validation = evaluate( model, protein, ligand, targets, split_indices["validation"], scaler, ) history.append( { "epoch": epoch, "train_standardized_mse": float(np.mean(losses)), "validation": validation, "learning_rate": optimizer.param_groups[0]["lr"], } ) if validation["rmse"] < best_validation: best_validation = validation["rmse"] best_epoch = epoch best_state = copy.deepcopy( {key: value.detach().cpu() for key, value in model.state_dict().items()} ) remaining_patience = args.patience else: remaining_patience -= 1 if remaining_patience <= 0: break model.load_state_dict(best_state) test_metrics = evaluate( model, protein, ligand, targets, split_indices["test"], scaler ) validation_metrics = evaluate( model, protein, ligand, targets, split_indices["validation"], scaler ) args.output_dir.mkdir(parents=True, exist_ok=True) save_file(best_state, args.output_dir / "model.safetensors") encoder_metadata_path = args.embeddings.with_suffix(".json") encoder_metadata = ( json.loads(encoder_metadata_path.read_text()) if encoder_metadata_path.exists() else None ) target_name = ( args.target_name or (encoder_metadata.get("target_name") if encoder_metadata else None) or "target" ) report = { "target": target_name, "embedding_file": str(args.embeddings), "manifest": str(args.manifest), "rows": {name: len(indices) for name, indices in split_indices.items()}, "seed": args.seed, "device": str(device), "best_epoch": best_epoch, "target_scaler": {"mean": scaler.mean, "std": scaler.std}, "head": { "protein_dim": int(protein.shape[1]), "ligand_dim": int(ligand.shape[1]), "projection_dim": args.projection_dim, "hidden_dim": args.hidden_dim, "dropout": args.dropout, }, "optimizer": { "name": "AdamW", "learning_rate": args.learning_rate, "weight_decay": args.weight_decay, "warmup_ratio": args.warmup_ratio, "warmup_steps": warmup_steps, "total_steps": total_steps, }, "validation": validation_metrics, "test": test_metrics, "encoders": encoder_metadata, "history": history, } (args.output_dir / "report.json").write_text(json.dumps(report, indent=2) + "\n") (args.output_dir / "config.json").write_text( json.dumps( { "format": "MitoInteract-v2-head", "target": target_name, "target_scaler": report["target_scaler"], "head": report["head"], "encoders": encoder_metadata, }, indent=2, ) + "\n" ) print( json.dumps( {key: value for key, value in report.items() if key != "history"}, indent=2 ) ) if __name__ == "__main__": main()