File size: 8,774 Bytes
8065faa | 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 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 | #!/usr/bin/env python
"""Evaluate a finetuned model on a labelled test set.
Three modes, one entry point:
# segmentation metrics (mIoU and Dice, per class)
python evaluate.py --config configs/aff_base_finetune_512_fpw.yaml
# foot-process-width geometry metrics from the paper
python evaluate.py --config <cfg> --mode fpw --out-json output/fpw.json
python evaluate.py --config <cfg> --mode fpw --seeds 42,77,2026
# side-by-side figure comparing several trained models
python evaluate.py --mode compare \
--config configs/aff_base_finetune_512_fpw.yaml \
--config configs/vit_base_finetune_fpn_512.yaml \
--label AFF-MAE --label MAE
"""
import argparse
import json
import logging
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from affmae.config import load_config # noqa: E402
from affmae.eval.fpw import ( # noqa: E402
FpwParams,
evaluate_fpw,
evaluate_fpw_across_seeds,
format_seed_summary,
json_safe,
parse_grid_size,
)
from affmae.eval.segmentation import ( # noqa: E402
compare_backends,
compare_models,
evaluate_segmentation,
)
from affmae.utils.dist import resolve_device # noqa: E402
from affmae.utils.env import load_dotenv # noqa: E402
from affmae.utils.misc import set_seed, setup_logging # noqa: E402
from affmae.utils.paths import output_path # noqa: E402
def _parse_boxes(raw):
"""Parse repeated ``x,y,w,h`` strings into tuples of int.
Args:
raw: sequence of comma-separated strings.
Returns:
List of (x, y, w, h) tuples.
Raises:
ValueError: if any entry does not have four parts.
"""
boxes = []
for item in raw:
parts = [int(value) for value in item.split(",")]
if len(parts) != 4:
raise ValueError(f"zoom box {item!r} must be x,y,w,h")
boxes.append(tuple(parts))
return boxes
def _fpw_params(args) -> FpwParams:
"""Build FpwParams from the parsed arguments."""
return FpwParams(
pgbmi_class=args.pgbmi_class,
slit_class=args.slit_class,
slit_min_area=args.slit_min_area,
slit_max_area=args.slit_max_area,
slit_circularity=args.slit_circularity,
pgbmi_dilate=args.pgbmi_dilate,
min_pgbmi_area=args.min_pgbmi_area,
min_segment_iou=args.min_segment_iou,
pixel_size_nm=args.pixel_size_nm,
eval_grid_size=parse_grid_size(args.eval_grid_size),
)
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=__doc__.split("\n")[0],
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=__doc__)
parser.add_argument("--config", action="append", required=True,
help="Training YAML. Repeat for --mode compare.")
parser.add_argument("--mode",
choices=("metrics", "fpw", "compare", "backends"),
default="metrics",
help="metrics: mIoU/Dice. fpw: paper geometry metrics. "
"compare: side-by-side figure. backends: score one "
"checkpoint through each kernel path, to see what "
"the faster fused decoder costs in accuracy.")
parser.add_argument("--checkpoint", action="append", default=None,
help="Explicit checkpoint; defaults to "
"<output_dir>/<name>/last_model.pth. Repeat for "
"--mode compare. May contain {seed}.")
parser.add_argument("--device", default=None, help="cuda | cpu | mps.")
parser.add_argument("--out-json", default=None,
help="Write the results here as JSON.")
parser.add_argument("--seed", type=int, default=42)
fpw = parser.add_argument_group("fpw mode")
fpw.add_argument("--seeds", default=None,
help="Comma-separated seeds, e.g. 42,77,2026. Each reads "
"<output_dir>/<name>_seed<seed>/last_model.pth.")
fpw.add_argument("--pgbmi-class", type=int, default=1)
fpw.add_argument("--slit-class", type=int, default=2)
fpw.add_argument("--slit-min-area", type=float, default=4.0)
fpw.add_argument("--slit-max-area", type=float, default=400.0)
fpw.add_argument("--slit-circularity", type=float, default=0.4)
fpw.add_argument("--pgbmi-dilate", type=int, default=3)
fpw.add_argument("--min-pgbmi-area", type=float, default=16.0)
fpw.add_argument("--min-segment-iou", type=float, default=0.1)
fpw.add_argument("--pixel-size-nm", type=float, default=1.0)
fpw.add_argument("--eval-grid-size", default="1024",
help="Reference grid for pixel distances: N or W,H.")
fpw.add_argument("--vis-dir", default=None,
help="Write per-image geometry overlays here.")
fpw.add_argument("--max-vis", type=int, default=None)
compare = parser.add_argument_group("compare mode")
compare.add_argument("--label", action="append", default=None,
help="Display name per config; defaults to model_type.")
compare.add_argument("--indices", type=int, nargs="+", default=[19, 91, 76],
help="Test-set sample indices to plot, one row each.")
compare.add_argument("--zoom-box", action="append", default=None,
metavar="X,Y,W,H",
help="Zoom window per plotted row.")
compare.add_argument("--out", default=None,
help="Figure path; defaults to "
"output/model_comparison.pdf.")
return parser
def main() -> None:
parser = build_parser()
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(message)s")
load_dotenv()
cfgs = [load_config(path) for path in args.config]
device = resolve_device(args.device or getattr(cfgs[0], "device", None))
for cfg in cfgs:
cfg.device = device
if args.mode != "compare" and len(cfgs) > 1:
parser.error(f"--mode {args.mode} takes one --config; use --mode compare")
checkpoints = args.checkpoint or [None] * len(cfgs)
if len(checkpoints) != len(cfgs):
parser.error(f"got {len(cfgs)} config(s) but {len(checkpoints)} "
f"--checkpoint value(s).")
if args.mode == "compare":
labels = args.label or [cfg.model_type for cfg in cfgs]
if len(labels) != len(cfgs):
parser.error(f"got {len(cfgs)} config(s) but {len(labels)} label(s).")
boxes = _parse_boxes(args.zoom_box) if args.zoom_box else None
if boxes and len(boxes) != len(args.indices):
parser.error(f"got {len(boxes)} zoom box(es) for "
f"{len(args.indices)} index/indices.")
set_seed(args.seed)
out = args.out or output_path("model_comparison.pdf", create_parent=True)
setup_logging(os.path.dirname(os.path.abspath(out)))
path = compare_models(cfgs, labels, out, indices=args.indices,
zoom_boxes=boxes, checkpoints=checkpoints)
logging.info("wrote %s", path)
return
cfg = cfgs[0]
exp_dir = os.path.join(cfg.output_dir, cfg.name)
os.makedirs(exp_dir, exist_ok=True)
setup_logging(exp_dir)
if args.mode == "backends":
results = compare_backends(cfg, checkpoints[0])
default_json = os.path.join(exp_dir, "backend_comparison.json")
elif args.mode == "fpw":
params = _fpw_params(args)
seeds = [int(s) for s in args.seeds.split(",") if s.strip()] \
if args.seeds else None
if seeds:
results = evaluate_fpw_across_seeds(
cfg, seeds, params, checkpoint=checkpoints[0],
vis_dir=args.vis_dir, max_vis=args.max_vis)
print(format_seed_summary(results["cross_seed_summary"]))
else:
results = evaluate_fpw(
cfg, params, seed=args.seed, checkpoint=checkpoints[0],
vis_dir=args.vis_dir, max_vis=args.max_vis)
print(json.dumps(json_safe(results["summary"]), indent=2))
results = json_safe(results)
default_json = os.path.join(exp_dir, "fpw_metrics.json")
else:
results = evaluate_segmentation(cfg, checkpoints[0])
default_json = None
out_json = args.out_json or default_json
if out_json:
os.makedirs(os.path.dirname(os.path.abspath(out_json)) or ".",
exist_ok=True)
with open(out_json, "w") as handle:
json.dump(results, handle, indent=2)
logging.info("wrote %s", out_json)
if __name__ == "__main__":
main()
|