Download code/training/src/training_validation/valid.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 20.2 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/valid.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/valid.py
-
curl -L -o valid.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/valid.py
20.2 kB
| 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 | |
| class Cluster: | |
| cluster_id: int | |
| raw_mask: np.ndarray | |
| mask: np.ndarray | |
| def pixel_count(self) -> int: | |
| return int(self.mask.sum()) | |
| def raw_pixel_count(self) -> int: | |
| return int(self.raw_mask.sum()) | |
| 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))) | |
| 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)) | |
| 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() | |