CD-Models / utils /unified_evaluator.py
Dineth Perera
Publish tested dataset winners and benchmark rankings
ce209f5
Raw
History Blame Contribute Delete
14.4 kB
from __future__ import annotations
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
import torch
from torch.utils.data import DataLoader, Subset
from datasets.cd_dataset import CDDataset
from utils.metrics import BinaryMetrics, BoundaryMetrics, normalize_binary_prediction
from utils.model_adapters import BaseModelAdapter
from utils.dataset_cache import dataloader_kwargs
from utils.profiling import GpuProfiler, ProfilingUnavailable, count_flops, count_parameters
from utils.qualitative import (
denormalize,
manifest_ids,
rank_for_sample,
safe_sample_id,
save_binary_prediction,
save_probability_map,
save_visual_panel,
select_or_load_manifest,
)
from utils.results_writer import append_to_comparison_table, save_metrics
ROOT = Path(__file__).resolve().parents[1]
def load_state_dict(checkpoint_path: Path) -> dict:
checkpoint = torch.load(checkpoint_path, map_location="cpu")
if isinstance(checkpoint, dict):
for key in ("model_state_dict", "state_dict", "model"):
if key in checkpoint and isinstance(checkpoint[key], dict):
return checkpoint[key]
if all(torch.is_tensor(v) for v in checkpoint.values()):
return checkpoint
raise RuntimeError(f"Checkpoint {checkpoint_path} does not contain a recognized PyTorch state_dict.")
def evaluate_torch_model(
*,
model_name: str,
dataset_cfg: dict,
model: torch.nn.Module,
checkpoint_path: Path,
forward_fn: Callable[[torch.nn.Module, torch.Tensor, torch.Tensor], torch.Tensor],
device: torch.device,
batch_size: int | None = None,
max_batches: int | None = None,
strict_profiling: bool = True,
output_dir: Path | None = None,
) -> tuple[dict, int]:
dataset_name = dataset_cfg["name"]
out_dir = output_dir or ROOT / "results" / model_name / dataset_name
pred_dir = out_dir / "predictions" / "test"
prob_dir = out_dir / "predictions" / "test_prob"
visual_dir = out_dir / "visuals" / "selected_20"
eval_cfg = dataset_cfg.get("eval", {})
threshold = float(eval_cfg.get("threshold", 0.5))
boundary_tolerance = int(eval_cfg.get("boundary_tolerance", 2))
state = load_state_dict(checkpoint_path)
model.load_state_dict(state, strict=True)
model.to(device)
model.eval()
ds = CDDataset(dataset_cfg["data_root"], "test", cfg=dataset_cfg, return_format="tuple")
if max_batches is not None:
ds_for_loader = Subset(ds, range(min(len(ds), max_batches * int(batch_size or dataset_cfg.get("batch_size", 1)))))
else:
ds_for_loader = ds
loader = DataLoader(
ds_for_loader,
batch_size=int(batch_size or dataset_cfg.get("batch_size", 8)),
shuffle=False,
**dataloader_kwargs(dataset_cfg, torch.cuda.is_available()),
)
manifest = select_or_load_manifest(dataset_cfg)
selected = manifest_ids(manifest)
metrics = BinaryMetrics(threshold=threshold)
boundary = BoundaryMetrics(tolerance=boundary_tolerance)
param_metrics = count_parameters(model)
flops_metrics: dict[str, object]
profiling_errors: list[str] = []
try:
image_size = int(dataset_cfg.get("img_size", 256))
flops_metrics = count_flops(
model,
lambda: (
torch.zeros(1, 3, image_size, image_size),
torch.zeros(1, 3, image_size, image_size),
),
device,
)
except ProfilingUnavailable as exc:
flops_metrics = {
"flops": None,
"flops_g": None,
"flops_input_shape": [[1, 3, int(dataset_cfg.get("img_size", 256)), int(dataset_cfg.get("img_size", 256))]] * 2,
"flops_library": None,
"flops_error": str(exc),
}
profiling_errors.append(str(exc))
warmup_batches = min(2, len(loader))
timed_images = 0
model_time = 0.0
end_to_end_start = time.perf_counter()
mean_a = dataset_cfg.get("mean_a", [0.485, 0.456, 0.406])
std_a = dataset_cfg.get("std_a", [0.229, 0.224, 0.225])
mean_b = dataset_cfg.get("mean_b", mean_a)
std_b = dataset_cfg.get("std_b", std_a)
with torch.inference_mode(), GpuProfiler(device=device, required=False) as gpu_profiler:
for batch_idx, (a, b, mask, names) in enumerate(loader):
a = a.to(device, non_blocking=True)
b = b.to(device, non_blocking=True)
mask_device = mask.to(device, non_blocking=True)
if device.type == "cuda":
torch.cuda.synchronize(device)
start = time.perf_counter()
outputs = forward_fn(model, a, b)
if device.type == "cuda":
torch.cuda.synchronize(device)
elapsed = time.perf_counter() - start
pred, prob = normalize_binary_prediction(outputs.detach().cpu(), threshold=threshold)
metrics.update(outputs.detach().cpu(), mask)
boundary.update(pred, mask)
if batch_idx >= warmup_batches:
model_time += elapsed
timed_images += int(a.shape[0])
for i, sample_id in enumerate(names):
clean_id = safe_sample_id(str(sample_id))
pred_i = pred[i]
save_binary_prediction(pred_i, pred_dir / f"{clean_id}_pred.png")
prob_i = prob[i] if prob is not None else None
if prob_i is not None:
save_probability_map(prob_i, prob_dir / f"{clean_id}_prob.png")
if str(sample_id) in selected:
rank = rank_for_sample(manifest, str(sample_id))
a_vis = denormalize(a[i].detach().cpu(), mean_a, std_a)
b_vis = denormalize(b[i].detach().cpu(), mean_b, std_b)
save_visual_panel(
a_vis,
b_vis,
mask[i],
pred_i,
visual_dir / f"{rank:02d}_{clean_id}_panel.png",
prob=prob_i,
)
del mask_device
end_to_end_time = time.perf_counter() - end_to_end_start
split_metrics = metrics.compute()
split_metrics.update(boundary.compute())
split_metrics.update(param_metrics)
split_metrics.update(flops_metrics)
split_metrics.update(gpu_profiler.summary())
if split_metrics.get("gpu_profiling_error"):
profiling_errors.append(str(split_metrics["gpu_profiling_error"]))
fps_model_only = timed_images / model_time if model_time > 0 else None
fps_end_to_end = len(ds_for_loader) / end_to_end_time if end_to_end_time > 0 else None
status = "complete" if not (strict_profiling and profiling_errors) else "incomplete"
split_metrics.update({
"model": model_name,
"dataset": dataset_name,
"split": "test",
"checkpoint": str(checkpoint_path),
"threshold": threshold,
"fps": fps_model_only,
"fps_model_only": fps_model_only,
"fps_end_to_end": fps_end_to_end,
"num_timed_images": timed_images,
"warmup_batches": warmup_batches,
"timing_device": str(device),
"test_num_samples": len(ds_for_loader),
"visual_sample_manifest": str(ROOT / "results" / "qualitative_samples" / dataset_name / "sample_manifest.json"),
"prediction_dir": str(pred_dir),
"visual_dir": str(visual_dir),
"timestamp": datetime.now(timezone.utc).isoformat(),
"status": status,
"profiling_errors": profiling_errors,
})
save_metrics(model_name, dataset_name, "test", split_metrics)
append_to_comparison_table()
return split_metrics, 0 if status == "complete" else 1
def evaluate_with_adapter(
*,
model_name: str,
dataset_cfg: dict,
model_config: dict,
adapter: BaseModelAdapter,
checkpoint_path: Path,
device: torch.device,
batch_size: int | None = None,
max_batches: int | None = None,
strict_profiling: bool = True,
output_dir: Path | None = None,
) -> tuple[dict, int]:
if not adapter.supports_inprocess_eval:
raise RuntimeError(f"{model_name} does not support in-process evaluation: {adapter.notes_or_failure_reason}")
dataset_name = dataset_cfg["name"]
out_dir = output_dir or ROOT / "results" / model_name / dataset_name
pred_dir = out_dir / "predictions" / "test"
prob_dir = out_dir / "predictions" / "test_prob"
visual_dir = out_dir / "visuals" / "selected_20"
eval_cfg = dataset_cfg.get("eval", {})
threshold = float(eval_cfg.get("threshold", 0.5))
boundary_tolerance = int(eval_cfg.get("boundary_tolerance", 2))
model = adapter.build_model(model_config, dataset_cfg, device)
adapter.load_checkpoint(model, checkpoint_path, device)
model.to(device)
model.eval()
ds = CDDataset(dataset_cfg["data_root"], "test", cfg=dataset_cfg, return_format="tuple")
if max_batches is not None:
ds_for_loader = Subset(ds, range(min(len(ds), max_batches * int(batch_size or dataset_cfg.get("batch_size", 1)))))
else:
ds_for_loader = ds
loader = DataLoader(
ds_for_loader,
batch_size=int(batch_size or dataset_cfg.get("batch_size", 8)),
shuffle=False,
**dataloader_kwargs(dataset_cfg, torch.cuda.is_available()),
)
manifest = select_or_load_manifest(dataset_cfg)
selected = manifest_ids(manifest)
metrics = BinaryMetrics(threshold=threshold)
boundary = BoundaryMetrics(tolerance=boundary_tolerance)
param_metrics = count_parameters(model)
profiling_errors: list[str] = []
try:
if not adapter.supports_flops:
raise ProfilingUnavailable(f"{model_name} adapter does not support FLOPs: {adapter.notes_or_failure_reason}")
flops_metrics = count_flops(model, lambda: adapter.get_dummy_inputs(dataset_cfg, device), device)
except ProfilingUnavailable as exc:
flops_metrics = {
"flops": None,
"flops_g": None,
"flops_input_shape": None,
"flops_library": None,
"flops_error": str(exc),
}
profiling_errors.append(str(exc))
warmup_batches = min(2, len(loader))
timed_images = 0
model_time = 0.0
end_to_end_start = time.perf_counter()
mean_a = dataset_cfg.get("mean_a", [0.485, 0.456, 0.406])
std_a = dataset_cfg.get("std_a", [0.229, 0.224, 0.225])
mean_b = dataset_cfg.get("mean_b", mean_a)
std_b = dataset_cfg.get("std_b", std_a)
with torch.inference_mode(), GpuProfiler(device=device, required=False) as gpu_profiler:
for batch_idx, batch in enumerate(loader):
if device.type == "cuda":
torch.cuda.synchronize(device)
start = time.perf_counter()
raw_output = adapter.forward(model, batch, device)
if device.type == "cuda":
torch.cuda.synchronize(device)
elapsed = time.perf_counter() - start
a, b, mask, names = batch
normalized = adapter.normalize_output(raw_output, batch, dataset_cfg)
metrics.update(normalized.metric_tensor, mask)
boundary.update(normalized.binary, mask)
if batch_idx >= warmup_batches:
model_time += elapsed
timed_images += int(a.shape[0])
for i, sample_id in enumerate(names):
clean_id = safe_sample_id(str(sample_id))
pred_i = normalized.binary[i]
save_binary_prediction(pred_i, pred_dir / f"{clean_id}_pred.png")
prob_i = normalized.score[i] if normalized.score is not None else None
if prob_i is not None:
save_probability_map(prob_i, prob_dir / f"{clean_id}_prob.png")
if str(sample_id) in selected:
rank = rank_for_sample(manifest, str(sample_id))
a_vis = denormalize(a[i].detach().cpu(), mean_a, std_a)
b_vis = denormalize(b[i].detach().cpu(), mean_b, std_b)
save_visual_panel(
a_vis,
b_vis,
mask[i],
pred_i,
visual_dir / f"{rank:02d}_{clean_id}_panel.png",
prob=prob_i,
)
end_to_end_time = time.perf_counter() - end_to_end_start
split_metrics = metrics.compute()
split_metrics.update(boundary.compute())
split_metrics.update(param_metrics)
split_metrics.update(flops_metrics)
split_metrics.update(gpu_profiler.summary())
if split_metrics.get("gpu_profiling_error"):
profiling_errors.append(str(split_metrics["gpu_profiling_error"]))
fps_model_only = timed_images / model_time if model_time > 0 else None
fps_end_to_end = len(ds_for_loader) / end_to_end_time if end_to_end_time > 0 else None
status = "complete" if not (strict_profiling and profiling_errors) else "incomplete"
split_metrics.update({
"model": model_name,
"dataset": dataset_name,
"split": "test",
"checkpoint": str(checkpoint_path),
"threshold": threshold,
"fps": fps_model_only,
"fps_model_only": fps_model_only,
"fps_end_to_end": fps_end_to_end,
"num_timed_images": timed_images,
"warmup_batches": warmup_batches,
"timing_device": str(device),
"test_num_samples": len(ds_for_loader),
"visual_sample_manifest": str(ROOT / "results" / "qualitative_samples" / dataset_name / "sample_manifest.json"),
"prediction_dir": str(pred_dir),
"visual_dir": str(visual_dir),
"timestamp": datetime.now(timezone.utc).isoformat(),
"status": status,
"profiling_errors": profiling_errors,
"adapter": {
"model_class_path": adapter.model_class_path,
"input_format": adapter.input_format,
"output_format": adapter.output_format,
"checkpoint_format": adapter.checkpoint_format,
"final_output_for_metrics": adapter.final_output_for_metrics,
"notes": adapter.notes_or_failure_reason,
},
})
save_metrics(model_name, dataset_name, "test", split_metrics)
append_to_comparison_table()
return split_metrics, 0 if status == "complete" else 1