File size: 3,364 Bytes
986e0b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
#!/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()