MitoInteract / recovery /scripts /train_head.py
Ethan Troy
feat: add Kd-only benchmark recovery track
38bce11
Raw
History Blame Contribute Delete
8.62 kB
#!/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()