lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
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
@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()