File size: 5,171 Bytes
fdae75f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
#!/usr/bin/env python
"""Score predicted ``.geff`` graphs against ground-truth ``.geff`` graphs.

Evaluates every dataset whose name appears in **both** ``--pred-dir`` and
``--gt-dir`` (others are skipped). No images are loaded: ground-truth tracks
come from the ``.geff`` and the voxel scale from the dataset's ``.zarr``
metadata (falling back to :data:`~voxmodel.io.DEFAULT_SCALE`).

Reports the run-level score — sample-size-weighted adjusted edge Jaccard plus
``0.1 x`` division Jaccard — via :func:`voxmodel.metrics.summarise`.

Usage:
    python scripts/evaluate.py --pred-dir out_geffs --gt-dir data/train
"""

from __future__ import annotations

import argparse
from pathlib import Path

import tracksdata as td
from geff import GeffMetadata

from voxmodel.io import DEFAULT_SCALE, open_dataset
from voxmodel.metrics import (
    evaluate as compute_metric,
    node_recall,
    per_sample_metrics,
    summarise,
)

from dataspec import DATASET_PATH


def _load_graph(geff_path: Path) -> td.graph.BaseGraph:
    result = td.graph.IndexedRXGraph.from_geff(geff_path)
    return result[0] if isinstance(result, tuple) else result


def _read_scale(gt_dir: Path, name: str) -> tuple[float, float, float]:
    """Return the dataset's (Z, Y, X) voxel scale, or DEFAULT_SCALE if no zarr.

    Only zarr metadata is read — the image itself is never loaded.
    """
    try:
        return open_dataset(gt_dir / name, load_image=False).scale
    except FileNotFoundError:
        return DEFAULT_SCALE


def _read_estimated_n_total(geff_path: Path) -> float:
    """Read ``estimated_number_of_nodes`` from a GEFF's metadata (NaN if absent)."""
    try:
        meta = GeffMetadata.read(geff_path)
    except Exception:
        return float("nan")
    val = (meta.extra or {}).get("estimated_number_of_nodes")
    return float(val) if val is not None else float("nan")


def evaluate_pairs(
    pred_dir: Path | str,
    gt_dir: Path | str,
    max_distance: float = 7.0,
) -> tuple[list[dict], list[str]]:
    """Score every dataset present in both dirs; return (per-sample rows, skipped).

    Rows are :func:`voxmodel.metrics.per_sample_metrics` dicts, ready for
    :func:`voxmodel.metrics.summarise`. No images are loaded.
    """
    pred_dir = Path(pred_dir)
    gt_dir = Path(gt_dir)

    pred_names = {p.stem for p in pred_dir.glob("*.geff")}
    gt_names = {p.stem for p in gt_dir.glob("*.geff")}
    names = sorted(pred_names & gt_names)
    print(f"{len(names)} datasets in both pred and GT (of {len(pred_names)} pred / {len(gt_names)} GT)")

    rows: list[dict] = []
    skipped: list[str] = []
    for name in names:
        gt_path = gt_dir / f"{name}.geff"
        try:
            pred_graph = _load_graph(pred_dir / f"{name}.geff")
            gt_graph = _load_graph(gt_path)
            scale = _read_scale(gt_dir, name)
            er = compute_metric(pred_graph, gt_graph, scale=scale, max_distance=max_distance)
            recall = (
                node_recall(pred_graph, gt_graph)
                if pred_graph.num_edges() > 0 and pred_graph.num_nodes() > 0
                else 0.0
            )
            n_total = _read_estimated_n_total(gt_path)
        except Exception as exc:  # unreadable/partial geff, etc.
            skipped.append(name)
            print(f"  SKIP {name}: {type(exc).__name__}: {exc}")
            continue

        rows.append(per_sample_metrics(er, n_total, recall))
        print(
            f"  {name}: edge TP/FP/FN={er.edge_tp}/{er.edge_fp}/{er.edge_fn} "
            f"div TP/FP/FN={er.division_tp}/{er.division_fp}/{er.division_fn} "
            f"n_pred={er.num_pred_nodes}"
        )

    if skipped:
        print(f"\nSkipped {len(skipped)} unreadable datasets: {skipped}")
    return rows, skipped


def evaluate_run(run: dict, max_distance: float = 7.0) -> list[dict]:
    """Score the predicted geffs in ``run['dir']`` against GT in ``DATASET_PATH``.

    Thin shim so ``scripts/predict_unet_transformer.py`` can evaluate a fresh
    prediction run right after inference.
    """
    rows, _ = evaluate_pairs(run["dir"], DATASET_PATH, max_distance=max_distance)
    return rows


def main() -> None:
    parser = argparse.ArgumentParser(description="Score predicted .geff graphs against ground-truth .geff graphs.")
    parser.add_argument("--pred-dir", type=Path, required=True, help="Directory of predicted .geff files.")
    parser.add_argument("--gt-dir", type=Path, required=True, help="Directory of ground-truth .geff files.")
    parser.add_argument("--max-distance", type=float, default=7.0)
    args = parser.parse_args()

    rows, _ = evaluate_pairs(args.pred_dir, args.gt_dir, max_distance=args.max_distance)
    s = summarise(rows)
    print("\n=== Summary ===")
    print(
        f"n={s['n']}  score={s['score']:.4f}  "
        f"edge_jaccard={s['edge_jaccard']:.4f}  "
        f"adj_edge_jaccard={s['adj_edge_jaccard']:.4f} (n_adj={s['n_adj']})  "
        f"division_jaccard={s['division_jaccard']:.4f} "
        f"(TP={s['division_tp']} FP={s['division_fp']} FN={s['division_fn']})  "
        f"node_recall={s['node_recall']:.4f}"
    )


if __name__ == "__main__":
    main()