from __future__ import annotations import argparse import concurrent.futures import glob import os import re import sys from dataclasses import dataclass from functools import lru_cache from pathlib import Path from typing import Any import numpy as np import pandas as pd import torch from scipy.ndimage import binary_dilation, label as cc_label from tqdm import tqdm try: from .common import ( build_dataloader, build_dataset, build_model, get_label, is_ram_chunk_dataset, load_config, load_model_checkpoint, maybe_copy_best, pack_inputs, shutdown_dataloader, write_json, ) from .logger import ExperimentLogger except ImportError: code_root = Path(__file__).resolve().parents[2] if str(code_root) not in sys.path: sys.path.insert(0, str(code_root)) from src.training_validation.common import ( # type: ignore build_dataloader, build_dataset, build_model, get_label, is_ram_chunk_dataset, load_config, load_model_checkpoint, maybe_copy_best, pack_inputs, shutdown_dataloader, write_json, ) from src.training_validation.logger import ExperimentLogger # type: ignore @dataclass class Cluster: cluster_id: int raw_mask: np.ndarray mask: np.ndarray @property def pixel_count(self) -> int: return int(self.mask.sum()) @property def raw_pixel_count(self) -> int: return int(self.raw_mask.sum()) @dataclass class PreparedTruth: clusters: list[Cluster] match_masks: list[np.ndarray] masks: list[np.ndarray] def _km_to_pixels(km: float, pixel_size_km: float) -> int: if km <= 0: return 0 return int(np.ceil(float(km) / float(pixel_size_km))) @lru_cache(maxsize=64) def _circular_footprint(radius: int) -> np.ndarray: if radius <= 0: return np.ones((1, 1), dtype=bool) y, x = np.ogrid[-radius : radius + 1, -radius : radius + 1] return (x * x + y * y) <= radius * radius def _dilate(mask: np.ndarray, radius: int) -> np.ndarray: if radius <= 0: return mask.astype(bool, copy=True) return binary_dilation(mask.astype(bool), structure=_circular_footprint(radius)) def _structure(connectivity: int) -> np.ndarray: if int(connectivity) == 4: return np.array([[0, 1, 0], [1, 1, 1], [0, 1, 0]], dtype=bool) return np.ones((3, 3), dtype=bool) def _components(mask: np.ndarray, min_pixels: int, connectivity: int) -> list[np.ndarray]: labeled, n_features = cc_label(mask.astype(bool), structure=_structure(connectivity)) out = [] for cid in range(1, n_features + 1): comp = labeled == cid if int(comp.sum()) >= int(min_pixels): out.append(comp) return out def cluster_field( field: np.ndarray, threshold: float, mode: str = "mode_like", min_cluster_pixels: int = 3, pixel_size_km: float = 2.0, merge_buffer_km: float = 12.0, cluster_mask_expansion_km: float = 0.0, connectivity: int = 8, ) -> list[Cluster]: valid_mask = np.isfinite(field) positive = (field >= float(threshold)) & valid_mask merge_radius = _km_to_pixels(float(merge_buffer_km), float(pixel_size_km)) expansion_radius = _km_to_pixels(float(cluster_mask_expansion_km), float(pixel_size_km)) mode = str(mode) if mode == "connected" or merge_radius <= 0: components = _components(positive, min_cluster_pixels, connectivity) clusters = [] for comp in components: final = _dilate(comp, expansion_radius) clusters.append(Cluster(len(clusters) + 1, comp, final)) return clusters if mode not in {"mode_like", "distance_merge"}: raise ValueError(f"unsupported cluster_mode: {mode}") support = _dilate(positive, merge_radius) support_labeled, n_support = cc_label(support, structure=_structure(connectivity)) clusters = [] for support_id in range(1, n_support + 1): support_mask = support_labeled == support_id raw_union = positive & support_mask if int(raw_union.sum()) < int(min_cluster_pixels): continue final = support_mask if mode == "mode_like" else _dilate(raw_union, expansion_radius) clusters.append(Cluster(len(clusters) + 1, raw_union, final)) return clusters def label_clusters(label_arr: np.ndarray, min_pixels: int = 1, connectivity: int = 8) -> list[Cluster]: comps = _components(label_arr > 0.5, min_pixels, connectivity) return [Cluster(i + 1, comp, comp) for i, comp in enumerate(comps)] def prepare_truth(target: np.ndarray, val_cfg: dict[str, Any]) -> PreparedTruth: pixel_size_km = float(val_cfg.get("pixel_size_km", 2.0)) label_buffer_px = _km_to_pixels(float(val_cfg.get("label_buffer_km", 0.0)), pixel_size_km) connectivity = int(val_cfg.get("connectivity", 8)) clusters = label_clusters( target, min_pixels=int(val_cfg.get("label_min_cluster_pixels", 1)), connectivity=connectivity, ) match_masks = [_dilate(truth.mask, label_buffer_px) for truth in clusters] masks = [truth.mask for truth in clusters] return PreparedTruth(clusters=clusters, match_masks=match_masks, masks=masks) def evaluate_scene_with_prepared_truth( pred: np.ndarray, prepared_truth: PreparedTruth, threshold: float, val_cfg: dict[str, Any], ) -> dict[str, int]: pixel_size_km = float(val_cfg.get("pixel_size_km", 2.0)) false_buffer_px = _km_to_pixels(float(val_cfg.get("model_false_buffer_km", 0.0)), pixel_size_km) connectivity = int(val_cfg.get("connectivity", 8)) pred_clusters = cluster_field( pred, threshold=threshold, mode=str(val_cfg.get("cluster_mode", "mode_like")), min_cluster_pixels=int(val_cfg.get("min_cluster_pixels", 3)), pixel_size_km=pixel_size_km, merge_buffer_km=float(val_cfg.get("merge_buffer_km", 12.0)), cluster_mask_expansion_km=float(val_cfg.get("cluster_mask_expansion_km", 0.0)), connectivity=connectivity, ) hits = 0 misses = 0 for truth_match in prepared_truth.match_masks: if any(bool(np.any(truth_match & pred_cluster.mask)) for pred_cluster in pred_clusters): hits += 1 else: misses += 1 falses = 0 for pred_cluster in pred_clusters: pred_match = _dilate(pred_cluster.mask, false_buffer_px) if not any(bool(np.any(pred_match & truth_mask)) for truth_mask in prepared_truth.masks): falses += 1 return { "hits": int(hits), "misses": int(misses), "falses": int(falses), "truth_clusters": int(len(prepared_truth.clusters)), "pred_clusters": int(len(pred_clusters)), } def evaluate_scene( pred: np.ndarray, target: np.ndarray, threshold: float, val_cfg: dict[str, Any], ) -> dict[str, int]: return evaluate_scene_with_prepared_truth(pred, prepare_truth(target, val_cfg), threshold, val_cfg) def scores(hits: int, misses: int, falses: int) -> dict[str, float | None]: pod = hits / (hits + misses) if hits + misses > 0 else None far = falses / (hits + falses) if hits + falses > 0 else None csi = hits / (hits + misses + falses) if hits + misses + falses > 0 else None f1 = 2 * hits / (2 * hits + misses + falses) if 2 * hits + misses + falses > 0 else None return {"POD": pod, "FAR": far, "CSI": csi, "F1": f1} def parse_epoch_from_path(path: Path) -> int: match = re.search(r"epoch_(\d+)", path.name) if match is None: return -1 return int(match.group(1)) @torch.no_grad() def evaluate_checkpoint( checkpoint_path: Path, config: dict[str, Any], dataset, loader, device: torch.device, input_sources: list[str], label_key: str, thresholds: list[float], truth_cache: dict[str, PreparedTruth] | None = None, ) -> list[dict[str, Any]]: model = build_model(config).to(device) payload = load_model_checkpoint(model, checkpoint_path, device) model.eval() epoch = int(payload.get("epoch", -1)) val_cfg = dict(config.get("validation", {})) eval_workers = int(val_cfg.get("eval_workers", min(32, max(1, (os.cpu_count() or 1) // 2)))) totals = { threshold: {"hits": 0, "misses": 0, "falses": 0, "truth_clusters": 0, "pred_clusters": 0} for threshold in thresholds } samples = 0 def _evaluate_one(scene_pred: np.ndarray, prepared_truth: PreparedTruth, threshold: float) -> tuple[float, dict[str, int]]: return threshold, evaluate_scene_with_prepared_truth(scene_pred, prepared_truth, threshold, val_cfg) def _truth_for_scene(sample_time: str | None, target: np.ndarray) -> PreparedTruth: if truth_cache is None or sample_time is None: return prepare_truth(target, val_cfg) prepared = truth_cache.get(sample_time) if prepared is None: prepared = prepare_truth(target, val_cfg) truth_cache[sample_time] = prepared return prepared def _consume_loader(active_loader, desc: str, executor: concurrent.futures.Executor | None) -> int: nonlocal samples chunk_samples = 0 for batch in tqdm(active_loader, desc=desc, dynamic_ncols=True): if batch is None: continue x = pack_inputs(batch, input_sources, device) y = get_label(batch, label_key, device) pred = model(x)["ci"].detach().cpu().numpy() truth = y.detach().cpu().numpy() batch_times = batch.get("time") pending = [] for i in range(pred.shape[0]): samples += 1 chunk_samples += 1 sample_time = str(batch_times[i]) if isinstance(batch_times, (list, tuple)) else None prepared_truth = _truth_for_scene(sample_time, truth[i]) for threshold in thresholds: if executor is None: scene = evaluate_scene_with_prepared_truth(pred[i], prepared_truth, threshold, val_cfg) for key, value in scene.items(): totals[threshold][key] += int(value) else: pending.append(executor.submit(_evaluate_one, pred[i], prepared_truth, threshold)) for future in concurrent.futures.as_completed(pending): threshold, scene = future.result() for key, value in scene.items(): totals[threshold][key] += int(value) return chunk_samples executor = None if eval_workers > 1: executor = concurrent.futures.ThreadPoolExecutor(max_workers=eval_workers) try: chunk_count = 1 if is_ram_chunk_dataset(dataset): chunk_count = int(dataset.num_chunks) dataset.load_chunk_sync(0, free_current_before_load=True) for chunk_id in range(chunk_count): active_loader = build_dataloader(config, dataset, mode="valid") iterator = iter(active_loader) next_chunk = chunk_id + 1 if next_chunk < chunk_count: dataset.start_preload(next_chunk) _consume_loader(iterator, f"valid {checkpoint_path.name} chunk {chunk_id + 1}/{chunk_count}", executor) shutdown_dataloader(active_loader) if next_chunk < chunk_count: if not dataset.wait_for_preload_and_swap(): dataset.load_chunk_sync(next_chunk, free_current_before_load=True) else: _consume_loader(loader, f"valid {checkpoint_path.name}", executor) finally: if executor is not None: executor.shutdown(wait=True) rows = [] for threshold in thresholds: total = totals[threshold] metric = scores(total["hits"], total["misses"], total["falses"]) rows.append( { "checkpoint": str(checkpoint_path), "epoch": epoch, "threshold": float(threshold), "samples": int(samples), "chunks": int(chunk_count), **total, **metric, } ) return rows def checkpoint_paths(config: dict[str, Any]) -> list[Path]: val_cfg = dict(config.get("validation", {})) pattern = val_cfg.get("checkpoint_glob") if pattern is None: out_dir = Path(config.get("output_dir", config.get("checkpoint_dir", "runs/default"))) pattern = str(out_dir / "checkpoints" / "epoch_*_model.pt") paths = [Path(p) for p in sorted(glob.glob(str(pattern)))] if not paths: raise FileNotFoundError(f"no checkpoints matched: {pattern}") start_epoch = val_cfg.get("start_epoch") if start_epoch is not None: min_epoch = int(start_epoch) paths = [path for path in paths if parse_epoch_from_path(path) >= min_epoch] if not paths: raise FileNotFoundError(f"no checkpoints matched start_epoch >= {min_epoch}: {pattern}") stride = int(val_cfg.get("checkpoint_stride", 1) or 1) if stride > 1: paths = paths[::stride] max_checkpoints = val_cfg.get("max_checkpoints") if max_checkpoints is not None: paths = paths[: int(max_checkpoints)] return paths def _threshold_key(value: Any) -> str: return f"{float(value):.8g}" def _completed_epochs_from_existing(metrics_path: Path, thresholds: list[float]) -> tuple[list[dict[str, Any]], set[int]]: if not metrics_path.exists(): return [], set() metrics_df = pd.read_csv(metrics_path) if metrics_df.empty: return [], set() required = {_threshold_key(v) for v in thresholds} completed: set[int] = set() for epoch, group in metrics_df.groupby("epoch"): present = {_threshold_key(v) for v in group["threshold"].tolist()} if required.issubset(present): completed.add(int(epoch)) return metrics_df.to_dict("records"), completed def _drop_epoch_rows(rows: list[dict[str, Any]], epoch: int) -> list[dict[str, Any]]: return [row for row in rows if int(row.get("epoch", -1)) != int(epoch)] def _sort_metric_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: return sorted(rows, key=lambda row: (int(row.get("epoch", -1)), float(row.get("threshold", 0.0)))) def main() -> None: parser = argparse.ArgumentParser(description="Lightweight checkpoint validation by CSI threshold sweep.") parser.add_argument("--config", required=True, help="Experiment YAML path") parser.add_argument("--device", default=None, help="Override device, e.g. cuda:0 or cpu") parser.add_argument("--output-dir", default=None, help="Override validation output directory") args = parser.parse_args() config = load_config(args.config) requested_device = str(args.device or config.get("device") or "auto") if requested_device == "auto": requested_device = "cuda" if torch.cuda.is_available() else "cpu" device = torch.device(requested_device) val_cfg = dict(config.get("validation", {})) if args.output_dir is not None: val_cfg["output_dir"] = str(Path(args.output_dir).resolve()) split_cfg = dict(config.get("valid", {})) input_sources = list(split_cfg.get("input_sources", config.get("input_sources", config.get("required_inputs", ["concat"])))) label_key = str(split_cfg.get("label_key", split_cfg.get("target_label", "ci"))) config.setdefault("valid", {}) config["valid"].setdefault("input_sources", input_sources) config["valid"].setdefault("required_labels", [label_key]) thresholds = [float(v) for v in val_cfg.get("thresholds", [round(x * 0.1, 1) for x in range(1, 10)])] dataset = build_dataset(config, split=str(split_cfg.get("split", "valid")), mode="valid") loader = build_dataloader(config, dataset, mode="valid") configured_out = val_cfg.get("output_dir") out_dir = Path(configured_out) if configured_out else Path(config.get("output_dir", "runs/default")) / "validation" out_dir.mkdir(parents=True, exist_ok=True) logger = ExperimentLogger(config, mode="valid") logger.start() try: metrics_path = out_dir / "epoch_threshold_metrics.csv" summary_path = out_dir / "epoch_summary.csv" best_path = out_dir / "best_checkpoint.json" resume_existing = bool(val_cfg.get("resume_existing", True)) all_rows, completed_epochs = _completed_epochs_from_existing(metrics_path, thresholds) if resume_existing else ([], set()) if completed_epochs: print(f"resuming validation: found {len(completed_epochs)} completed epochs in {metrics_path}") def save_partial_results() -> dict[str, Any] | None: if not all_rows: return None metrics_df = pd.DataFrame(_sort_metric_rows(all_rows)) metrics_df.to_csv(metrics_path, index=False) valid_csi = metrics_df["CSI"].fillna(-1.0) best_idx = int(valid_csi.idxmax()) best_row = metrics_df.loc[best_idx].to_dict() summary_df = ( metrics_df.sort_values(["epoch", "CSI"], ascending=[True, False]) .groupby("checkpoint", as_index=False) .head(1) .sort_values("CSI", ascending=False) ) summary_df.to_csv(summary_path, index=False) best_payload = {"best": best_row, "metrics_path": str(metrics_path)} write_json(best_path, best_payload) maybe_copy_best( Path(str(best_row["checkpoint"])), out_dir / "best_model.pt", enabled=bool(val_cfg.get("copy_best_model", True)), ) return best_row truth_cache: dict[str, PreparedTruth] | None = {} if bool(val_cfg.get("cache_truth", True)) else None for ckpt in checkpoint_paths(config): ckpt_epoch = parse_epoch_from_path(ckpt) if ckpt_epoch in completed_epochs: print(f"skip completed checkpoint: {ckpt.name} epoch={ckpt_epoch}") continue all_rows = _drop_epoch_rows(all_rows, ckpt_epoch) rows = evaluate_checkpoint(ckpt, config, dataset, loader, device, input_sources, label_key, thresholds, truth_cache=truth_cache) all_rows.extend(rows) completed_epochs.add(int(rows[0]["epoch"]) if rows else ckpt_epoch) ckpt_df = pd.DataFrame(rows) if not ckpt_df.empty: best_ckpt_row = ckpt_df.loc[int(ckpt_df["CSI"].fillna(-1.0).idxmax())].to_dict() logger.log(best_ckpt_row, step=int(best_ckpt_row.get("epoch", 0)), prefix="valid_checkpoint") best_row = save_partial_results() if best_row is not None: print( "saved validation results:", metrics_path, "current_best_epoch=", best_row["epoch"], "threshold=", best_row["threshold"], "CSI=", best_row["CSI"], ) best_row = save_partial_results() if best_row is None: raise RuntimeError("validation produced no metric rows") logger.log(best_row, step=int(best_row.get("epoch", 0)), prefix="valid_best") logger.log_file(metrics_path, name="validation_threshold_metrics") logger.log_file(summary_path, name="validation_epoch_summary") logger.log_file(best_path, name="validation_best_checkpoint") print( "best checkpoint:", best_row["checkpoint"], "epoch=", best_row["epoch"], "threshold=", best_row["threshold"], "CSI=", best_row["CSI"], ) finally: shutdown_dataloader(loader) if is_ram_chunk_dataset(dataset): dataset.shutdown_preload() logger.finish() if __name__ == "__main__": main()