Saluki / scripts /predict.py
wuxing0105's picture
Upload folder using huggingface_hub (part 2)
986e0b8 verified
Raw
History Blame Contribute Delete
3.36 kB
#!/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()