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