#!/usr/bin/env python3 """Run Saluki prediction/evaluation without requiring targets.txt metadata.""" from __future__ import annotations import argparse import json from pathlib import Path from _bootstrap import configure_runtime ROOT = configure_runtime() def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Predict mRNA half-life from an official Saluki TFRecord split." ) parser.add_argument("data_dir", type=Path) parser.add_argument("weight", type=Path) parser.add_argument("--params", type=Path, default=ROOT / "conf" / "params.json") parser.add_argument("--head", type=int, default=0) parser.add_argument("--split", default="test") parser.add_argument("--batch-size", type=int, default=None) parser.add_argument("--out-dir", type=Path, default=ROOT / "output" / "predict") return parser.parse_args() def main() -> None: args = parse_args() if not args.params.is_file(): raise FileNotFoundError(args.params) if not args.weight.is_file(): raise FileNotFoundError(args.weight) if not (args.data_dir / "statistics.json").is_file(): raise FileNotFoundError(args.data_dir / "statistics.json") import h5py import numpy as np from scipy.stats import pearsonr from model.data import RnaDataset from model.saluki import SalukiModel, load_params params = load_params(args.params) batch_size = args.batch_size or int(params["train"]["batch_size"]) data = RnaDataset( str(args.data_dir), split_label=args.split, batch_size=batch_size, mode="eval" ) model = SalukiModel(params, head=args.head).restore(args.weight) predictions = np.asarray(model.predict(data.dataset, verbose=1), dtype=np.float32) targets = np.asarray(data.numpy(return_inputs=False), dtype=np.float32) targets = targets.reshape(predictions.shape) finite = bool(np.isfinite(predictions).all()) mse = float(np.mean(np.square(targets - predictions))) pred_flat = predictions.reshape(-1) target_flat = targets.reshape(-1) pearson = ( float(pearsonr(target_flat, pred_flat).statistic) if pred_flat.size > 1 else None ) denominator = float(np.sum(np.square(target_flat - target_flat.mean()))) r2 = ( float(1.0 - np.sum(np.square(target_flat - pred_flat)) / denominator) if denominator > 0 else None ) args.out_dir.mkdir(parents=True, exist_ok=True) with h5py.File(args.out_dir / "predictions.h5", "w") as handle: handle.create_dataset("predictions", data=predictions) handle.create_dataset("targets", data=targets) metrics = { "split": args.split, "head": args.head, "samples": int(predictions.shape[0]), "output_shape": list(predictions.shape), "dtype": str(predictions.dtype), "finite": finite, "nan_count": int(np.isnan(predictions).sum()), "inf_count": int(np.isinf(predictions).sum()), "mse": mse, "pearson_r": pearson, "r2": r2 } (args.out_dir / "metrics.json").write_text( json.dumps(metrics, indent=2) + "\n", encoding="utf-8" ) print(json.dumps(metrics, indent=2)) if not finite: raise RuntimeError("Predictions contain NaN or Inf") if __name__ == "__main__": main()