Download code/validation/src/run_threshold_sweep.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 15.9 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/run_threshold_sweep.py
- Command line
-
hf download hf://lsh9034/ci-net/code/validation/src/run_threshold_sweep.py
-
curl -L -o run_threshold_sweep.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/run_threshold_sweep.py
15.9 kB
| #!/usr/bin/env python3 | |
| """Run object-based validation for several model thresholds.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import os | |
| from copy import deepcopy | |
| from pathlib import Path | |
| from typing import Any | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| from .clusterers import create_clusterer | |
| from .loaders import CloudLabelLoader, ValidationJsonLoader, create_prediction_provider | |
| from .metrics import compute_scores | |
| from .threshold_sweep import ThresholdSweepValidator | |
| from .utils import expand_modes | |
| from .config import load_config as load_release_config | |
| DEFAULT_THRESHOLDS = [round(0.1 * idx, 1) for idx in range(1, 10)] | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Threshold-sweep object-based validation") | |
| parser.add_argument("--config", default="config_threshold_sweep.yaml", help="Path to YAML config") | |
| parser.add_argument("--data-source", choices=["Model"], default=None) | |
| parser.add_argument("--device", default=None, help="Accepted for a common CLI; validation runs on CPU") | |
| parser.add_argument("--dates", nargs="*", default=None, help="Override dates, e.g. 2024 2025") | |
| parser.add_argument("--target-filters", nargs="*", choices=["true_only", "all", "both"], default=None) | |
| parser.add_argument("--leadtime-modes", nargs="*", choices=["exact", "accumulate", "both"], default=None) | |
| parser.add_argument("--thresholds", nargs="*", type=float, default=None, help="Override thresholds, e.g. 0.1 0.3 0.5") | |
| parser.add_argument("--max-cases", type=int, default=None, help="Override validation.max_cases") | |
| parser.add_argument("--output-dir", default=None, help="Override paths.output_dir") | |
| return parser.parse_args() | |
| def apply_overrides(config: dict[str, Any], args: argparse.Namespace) -> dict[str, Any]: | |
| config = deepcopy(config) | |
| config.setdefault("validation", {}) | |
| config.setdefault("paths", {}) | |
| config.setdefault("threshold_sweep", {}) | |
| if args.data_source is not None: | |
| config["data_source"] = args.data_source | |
| if args.dates is not None and len(args.dates) > 0: | |
| config["dates"] = args.dates | |
| if args.target_filters is not None and len(args.target_filters) > 0: | |
| config["validation"]["target_filters"] = args.target_filters | |
| if args.leadtime_modes is not None and len(args.leadtime_modes) > 0: | |
| config["validation"]["leadtime_modes"] = args.leadtime_modes | |
| if args.thresholds is not None and len(args.thresholds) > 0: | |
| config["threshold_sweep"]["thresholds"] = args.thresholds | |
| if args.max_cases is not None: | |
| config["validation"]["max_cases"] = args.max_cases | |
| if args.output_dir is not None: | |
| config["paths"]["output_dir"] = args.output_dir | |
| return config | |
| def json_default(obj: Any) -> Any: | |
| if hasattr(obj, "item"): | |
| return obj.item() | |
| return str(obj) | |
| def write_json(data: Any, path: Path) -> None: | |
| with open(path, "w", encoding="utf-8") as f: | |
| json.dump(data, f, indent=2, ensure_ascii=False, default=json_default) | |
| def csv_value(value: Any) -> Any: | |
| if isinstance(value, list): | |
| return ";".join(str(item) for item in value) | |
| if isinstance(value, dict): | |
| return json.dumps(value, ensure_ascii=False) | |
| return value | |
| def write_csv(records: list[dict[str, Any]], path: Path) -> None: | |
| if not records: | |
| path.write_text("", encoding="utf-8") | |
| return | |
| fieldnames: list[str] = [] | |
| for record in records: | |
| for key in record: | |
| if key not in fieldnames: | |
| fieldnames.append(key) | |
| with open(path, "w", encoding="utf-8", newline="") as f: | |
| writer = csv.DictWriter(f, fieldnames=fieldnames) | |
| writer.writeheader() | |
| for record in records: | |
| writer.writerow({key: csv_value(record.get(key)) for key in fieldnames}) | |
| def threshold_values(config: dict[str, Any]) -> list[float]: | |
| values = (config.get("threshold_sweep") or {}).get("thresholds", DEFAULT_THRESHOLDS) | |
| thresholds = sorted({round(float(value), 10) for value in values}) | |
| if not thresholds: | |
| raise ValueError("No thresholds configured") | |
| return thresholds | |
| def threshold_dir_name(threshold: float) -> str: | |
| return f"threshold_{threshold:g}".replace(".", "p") | |
| def release_relative_path(value: str | Path, config: dict[str, Any]) -> str: | |
| path = Path(value) | |
| config_path = Path(config["_config_path"]) | |
| repository_root = config_path.parents[3] | |
| try: | |
| return path.resolve().relative_to(repository_root).as_posix() | |
| except ValueError: | |
| return path.name | |
| def create_threshold_clusterers(config: dict[str, Any], thresholds: list[float]): | |
| clusterers = {} | |
| for threshold in thresholds: | |
| threshold_config = deepcopy(config) | |
| threshold_config.setdefault("clusterer", {}) | |
| threshold_config["clusterer"]["threshold"] = float(threshold) | |
| clusterers[float(threshold)] = create_clusterer(threshold_config) | |
| return clusterers | |
| def combo_output_dir(config: dict[str, Any], date_str: str, target_filter: str, leadtime_mode: str) -> Path: | |
| del target_filter, leadtime_mode | |
| return Path(config["paths"]["output_dir"]) / str(config["data_source"]) / str(date_str) | |
| def summary_csv_row(threshold: float, summary: dict[str, Any]) -> dict[str, Any]: | |
| total = summary["total"] | |
| return { | |
| "threshold": float(threshold), | |
| "hits": total["hits"], | |
| "misses": total["misses"], | |
| "falses": total["falses"], | |
| "POD": total["POD"], | |
| "FAR": total["FAR"], | |
| "F1": total["F1"], | |
| "CSI": total["CSI"], | |
| "valid_labels": total["valid_labels"], | |
| "impossible": total["impossible"], | |
| "total_input_labels": total["total_input_labels"], | |
| "model_hit_clusters": total["model_hit_clusters"], | |
| "false_model_clusters": total["false_model_clusters"], | |
| "total_model_clusters": total["total_model_clusters"], | |
| } | |
| def apply_report_leadtime_range( | |
| summary: dict[str, Any], | |
| label_records: list[dict[str, Any]], | |
| model_cluster_records: list[dict[str, Any]], | |
| report_min: int, | |
| report_max: int, | |
| ) -> dict[str, Any]: | |
| """Report a lead-time subset after accumulation over the full evaluation range.""" | |
| selected_leadtimes = { | |
| int(leadtime): values | |
| for leadtime, values in summary["leadtime"].items() | |
| if report_min <= int(leadtime) <= report_max | |
| } | |
| if not selected_leadtimes: | |
| raise ValueError(f"No lead-time results in reporting range {report_min}..{report_max}") | |
| hits = sum(int(values["hits"]) for values in selected_leadtimes.values()) | |
| misses = sum(int(values["misses"]) for values in selected_leadtimes.values()) | |
| falses = sum(int(values["falses"]) for values in selected_leadtimes.values()) | |
| impossible = sum(int(values["impossible"]) for values in selected_leadtimes.values()) | |
| eligible_cloud_ids = { | |
| str(record["cloud_id"]) | |
| for record in label_records | |
| if report_min <= int(record["leadtime"]) <= report_max | |
| } | |
| model_hits = sum( | |
| 1 | |
| for record in model_cluster_records | |
| if not record["is_false"] | |
| and any(str(cloud_id) in eligible_cloud_ids for cloud_id in record.get("matched_cloud_ids", [])) | |
| ) | |
| report_total = compute_scores(hits, misses, falses) | |
| report_total.update( | |
| { | |
| "valid_labels": int(hits + misses), | |
| "impossible": int(impossible), | |
| "total_input_labels": int(hits + misses + impossible), | |
| "model_hit_clusters": int(model_hits), | |
| "false_model_clusters": int(falses), | |
| "total_model_clusters": int(model_hits + falses), | |
| } | |
| ) | |
| summary["evaluation_total"] = summary["total"] | |
| summary["total"] = report_total | |
| summary["leadtime"] = selected_leadtimes | |
| summary["report_leadtime_range"] = {"min": int(report_min), "max": int(report_max)} | |
| return summary | |
| def plot_pod_far(rows: list[dict[str, Any]], output_path: Path) -> None: | |
| fig, ax = plt.subplots(figsize=(6.5, 5.5), constrained_layout=True) | |
| valid_rows = [row for row in rows if row.get("POD") is not None and row.get("FAR") is not None] | |
| if valid_rows: | |
| far = [float(row["FAR"]) for row in valid_rows] | |
| pod = [float(row["POD"]) for row in valid_rows] | |
| thresholds = [float(row["threshold"]) for row in valid_rows] | |
| ax.plot(far, pod, marker="o", linewidth=1.8) | |
| for x, y, threshold in zip(far, pod, thresholds): | |
| ax.annotate(f"{threshold:g}", (x, y), textcoords="offset points", xytext=(5, 5), fontsize=8) | |
| else: | |
| ax.text(0.5, 0.5, "No valid POD/FAR points", ha="center", va="center", transform=ax.transAxes) | |
| ax.set_xlabel("FAR") | |
| ax.set_ylabel("POD") | |
| ax.set_title("POD-FAR Threshold Curve") | |
| ax.set_xlim(-0.02, 1.02) | |
| ax.set_ylim(-0.02, 1.02) | |
| ax.grid(True, linestyle=":", linewidth=0.6, alpha=0.7) | |
| fig.savefig(output_path, dpi=150) | |
| plt.close(fig) | |
| def write_details( | |
| out_dir: Path, | |
| threshold: float, | |
| label_records: list[dict[str, Any]], | |
| model_cluster_records: list[dict[str, Any]], | |
| missing_predictions: list[dict[str, Any]], | |
| ) -> None: | |
| detail_dir = out_dir / "details" / threshold_dir_name(threshold) | |
| detail_dir.mkdir(parents=True, exist_ok=True) | |
| write_json(label_records, detail_dir / "label_results.json") | |
| write_json(model_cluster_records, detail_dir / "model_cluster_results.json") | |
| write_json(missing_predictions, detail_dir / "missing_predictions.json") | |
| write_csv(label_records, detail_dir / "label_results.csv") | |
| write_csv(model_cluster_records, detail_dir / "model_cluster_results.csv") | |
| def run_for_date_and_filter(config: dict[str, Any], date_str: str, target_filter: str) -> dict[str, Any]: | |
| validation_config = config["validation"] | |
| sweep_config = config.get("threshold_sweep") or {} | |
| json_path = config["paths"]["validation_json_template"].format(date=date_str) | |
| thresholds = threshold_values(config) | |
| targets = ValidationJsonLoader(json_path).load_targets( | |
| target_filter=target_filter, | |
| leadtime_min=int(validation_config.get("leadtime_min", 10)), | |
| leadtime_max=int(validation_config.get("leadtime_max", 120)), | |
| max_cases=validation_config.get("max_cases"), | |
| ) | |
| label_loader = CloudLabelLoader(config["paths"]["temporal_overlapping_dir"]) | |
| provider = create_prediction_provider(config) | |
| clusterers = create_threshold_clusterers(config, thresholds) | |
| matching_config = config.get("matching") or {} | |
| performance_config = config.get("performance") or {} | |
| validator = ThresholdSweepValidator( | |
| targets=targets, | |
| label_loader=label_loader, | |
| prediction_provider=provider, | |
| clusterers=clusterers, | |
| pixel_size_km=float(matching_config.get("pixel_size_km", 2.0)), | |
| label_buffer_km=float(matching_config.get("label_buffer_km", 0.0)), | |
| model_false_buffer_km=float(matching_config.get("model_false_buffer_km", 0.0)), | |
| leadtime_min=int(validation_config.get("leadtime_min", 10)), | |
| leadtime_max=int(validation_config.get("leadtime_max", 120)), | |
| time_step=int(validation_config.get("time_step", 10)), | |
| buffer_backend=str(performance_config.get("buffer_backend", "auto")), | |
| ) | |
| sweep_result = validator.evaluate_raw() | |
| leadtime_modes = expand_modes(validation_config.get("leadtime_modes", ["exact"]), ("exact", "accumulate")) | |
| if len(leadtime_modes) != 1: | |
| raise ValueError("The public output layout requires exactly one lead-time mode") | |
| save_details = bool(sweep_config.get("save_details", False)) | |
| mode_summaries: dict[str, Any] = {} | |
| for leadtime_mode in leadtime_modes: | |
| out_dir = combo_output_dir(config, date_str, target_filter, leadtime_mode) | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| summary_items: list[dict[str, Any]] = [] | |
| csv_rows: list[dict[str, Any]] = [] | |
| for threshold in thresholds: | |
| raw_result = sweep_result.raw_results[threshold] | |
| label_records = validator.apply_leadtime_mode(raw_result.label_records, leadtime_mode) | |
| summary = validator.summarize(label_records, raw_result.model_cluster_records) | |
| report_min = int( | |
| validation_config.get("report_leadtime_min", validation_config.get("leadtime_min", 10)) | |
| ) | |
| report_max = int( | |
| validation_config.get("report_leadtime_max", validation_config.get("leadtime_max", 120)) | |
| ) | |
| summary = apply_report_leadtime_range( | |
| summary, | |
| label_records, | |
| raw_result.model_cluster_records, | |
| report_min, | |
| report_max, | |
| ) | |
| summary.update( | |
| { | |
| "date": date_str, | |
| "data_source": config["data_source"], | |
| "target_filter": target_filter, | |
| "leadtime_mode": leadtime_mode, | |
| "threshold": float(threshold), | |
| "validation_json_path": release_relative_path(json_path, config), | |
| "num_loaded_targets": len(targets), | |
| "num_report_targets": summary["total"]["total_input_labels"], | |
| "clusterer": dict(config.get("clusterer", {}), threshold=float(threshold)), | |
| "matching": config.get("matching", {}), | |
| } | |
| ) | |
| summary_items.append(summary) | |
| csv_rows.append(summary_csv_row(threshold, summary)) | |
| if save_details: | |
| write_details( | |
| out_dir=out_dir, | |
| threshold=threshold, | |
| label_records=label_records, | |
| model_cluster_records=raw_result.model_cluster_records, | |
| missing_predictions=raw_result.missing_predictions, | |
| ) | |
| total = summary["total"] | |
| print( | |
| f"[{date_str}][{config['data_source']}][{target_filter}][{leadtime_mode}] " | |
| f"threshold={threshold:g} POD={total['POD']} FAR={total['FAR']} " | |
| f"F1={total['F1']} CSI={total['CSI']} H={total['hits']} " | |
| f"M={total['misses']} F={total['falses']} impossible={total['impossible']}" | |
| ) | |
| payload = { | |
| "date": date_str, | |
| "data_source": config["data_source"], | |
| "target_filter": target_filter, | |
| "leadtime_mode": leadtime_mode, | |
| "thresholds": thresholds, | |
| "summaries": summary_items, | |
| } | |
| write_json(payload, out_dir / "threshold_summary.json") | |
| write_csv(csv_rows, out_dir / "threshold_summary.csv") | |
| plot_pod_far(csv_rows, out_dir / "pod_far_curve.png") | |
| mode_summaries[leadtime_mode] = payload | |
| return mode_summaries | |
| def main() -> None: | |
| args = parse_args() | |
| config_path = Path(args.config) | |
| if not config_path.is_absolute(): | |
| config_path = Path(os.getcwd()) / config_path | |
| config = apply_overrides(load_release_config(config_path), args) | |
| target_filters = expand_modes( | |
| config["validation"].get("target_filters", ["true_only"]), | |
| ("true_only", "all"), | |
| ) | |
| if len(target_filters) != 1: | |
| raise ValueError("The public output layout requires exactly one target filter") | |
| dates = [str(date) for date in config.get("dates", [])] | |
| if not dates: | |
| raise ValueError("No dates configured") | |
| for date_str in dates: | |
| for target_filter in target_filters: | |
| run_for_date_and_filter(config, date_str, target_filter) | |
| output_root = Path(config["paths"]["output_dir"]) / str(config["data_source"]) | |
| print(f"Validation summaries saved under {output_root}") | |
| if __name__ == "__main__": | |
| main() | |