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