cad-benchmark / validator.py
Gotech's picture
Deploy CAD Benchmark + Live 3D LC-KSVD Abnormality Detector & Slicer
c8af385 verified
Raw History Blame Contribute Delete
16.4 kB
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, Optional, Any, Tuple
import numpy as np
import nibabel as nib
try:
from schemas import ClassificationSample, LocalizationSample, ExplainabilitySample
except ImportError:
from .schemas import ClassificationSample, LocalizationSample, ExplainabilitySample
@dataclass
class ValidationResult:
is_valid: bool
errors: List[str] = field(default_factory=list)
model_name: Optional[str] = None
model_type: Optional[str] = None
classification_samples: Optional[List[ClassificationSample]] = None
localization_samples: Optional[List[LocalizationSample]] = None
explainability_samples: Optional[List[ExplainabilitySample]] = None
summary: Optional[Dict[str, Any]] = None
def _get_axcodes(img: nib.spatialimages.SpatialImage) -> Tuple[str, ...]:
try:
return nib.aff2axcodes(img.affine)
except Exception:
return ()
def _format_axcodes(axcodes: Tuple[str, ...]) -> str:
if not axcodes:
return "UNKNOWN"
return "".join(axcodes) + "+"
def _load_nifti_data_and_meta(
file_path: Path,
) -> Tuple[Optional[np.ndarray], Optional[Tuple[int, ...]], Optional[Tuple[float, ...]], Optional[Tuple[str, ...]], Optional[str]]:
try:
img = nib.load(str(file_path))
data = img.get_fdata(dtype=np.float32)
shape = tuple(img.shape)
spacing = tuple(float(z) for z in img.header.get_zooms()[:3])
axcodes = _get_axcodes(img)
return data, shape, spacing, axcodes, None
except Exception as e:
return None, None, None, None, str(e)
def validate_submission(
manifest: Dict[str, Any],
gt_index: Dict[str, Any],
files_dir: Optional[Path] = None,
) -> ValidationResult:
errors: List[str] = []
if not isinstance(manifest, dict):
return ValidationResult(is_valid=False, errors=["Manifest must be a JSON object / Python dict."])
model_name = manifest.get("model_name")
if not model_name or not isinstance(model_name, str) or not model_name.strip():
errors.append("Field 'model_name' must be a non-empty string.")
model_type = manifest.get("model_type")
valid_types = {"classification", "segmentation", "unified"}
if model_type not in valid_types:
errors.append(
f"Field 'model_type' must be one of {sorted(valid_types)}, got: {repr(model_type)}"
)
return ValidationResult(is_valid=False, errors=errors)
gt_cases = gt_index.get("cases", {})
if not gt_cases:
return ValidationResult(
is_valid=False,
errors=["Ground truth index contains no held-out cases. Cannot validate."],
)
classification_rows = []
segmentation_rows = []
if model_type == "classification":
if "segmentation" in manifest:
errors.append("Manifest declared model_type='classification' but contains a 'segmentation' block.")
if "predictions" not in manifest or not isinstance(manifest["predictions"], list):
errors.append("Classification manifest must contain a 'predictions' list.")
else:
classification_rows = manifest["predictions"]
elif model_type == "segmentation":
if "classification" in manifest:
errors.append("Manifest declared model_type='segmentation' but contains a 'classification' block.")
if "predictions" not in manifest or not isinstance(manifest["predictions"], list):
errors.append("Segmentation manifest must contain a 'predictions' list.")
else:
segmentation_rows = manifest["predictions"]
elif model_type == "unified":
has_cls = "classification" in manifest and isinstance(manifest["classification"], dict)
has_seg = "segmentation" in manifest and isinstance(manifest["segmentation"], dict)
if not has_cls or not has_seg:
missing = []
if not has_cls:
missing.append("'classification'")
if not has_seg:
missing.append("'segmentation'")
errors.append(
f"Manifest declared model_type='unified' but missing required block(s): {', '.join(missing)}."
)
else:
cls_preds = manifest["classification"].get("predictions")
seg_preds = manifest["segmentation"].get("predictions")
if not isinstance(cls_preds, list):
errors.append("Unified manifest 'classification' block must contain a 'predictions' list.")
else:
classification_rows = cls_preds
if not isinstance(seg_preds, list):
errors.append("Unified manifest 'segmentation' block must contain a 'predictions' list.")
else:
segmentation_rows = seg_preds
if errors:
return ValidationResult(is_valid=False, errors=errors, model_name=model_name, model_type=model_type)
seen_cls_scans = set()
cls_predictions_by_scan: Dict[str, Dict[str, int]] = {}
for idx, row in enumerate(classification_rows):
prefix = f"Classification prediction [{idx}]"
if not isinstance(row, dict):
errors.append(f"{prefix} must be a JSON object.")
continue
scan_id = row.get("scan_id")
if not scan_id or not isinstance(scan_id, str):
errors.append(f"{prefix} missing valid string 'scan_id'.")
continue
if scan_id not in gt_cases:
errors.append(f"{prefix} scan_id '{scan_id}' not found in held-out benchmark case list.")
continue
if scan_id in seen_cls_scans:
errors.append(f"{prefix} duplicate scan_id '{scan_id}' in classification predictions.")
continue
seen_cls_scans.add(scan_id)
expected_keys = {"scan_id", "2c", "2d"}
actual_keys = set(row.keys())
if actual_keys != expected_keys:
extra = actual_keys - expected_keys
missing = expected_keys - actual_keys
msg = f"{prefix} (scan_id '{scan_id}') keys mismatch."
if extra:
msg += f" Disallowed extra keys: {sorted(extra)}."
if missing:
msg += f" Missing required keys: {sorted(missing)}."
errors.append(msg)
continue
valid_row_values = True
for k in ["2c", "2d"]:
val = row[k]
if type(val) is not int or val not in (0, 1):
errors.append(
f"{prefix} (scan_id '{scan_id}') value for '{k}' must be int 0 or 1, got {type(val).__name__} ({repr(val)})."
)
valid_row_values = False
if valid_row_values:
cls_predictions_by_scan[scan_id] = {"2c": row["2c"], "2d": row["2d"]}
seen_finding_ids = set()
parsed_loc_items = []
parsed_exp_items = []
for idx, row in enumerate(segmentation_rows):
prefix = f"Segmentation prediction [{idx}]"
if not isinstance(row, dict):
errors.append(f"{prefix} must be a JSON object.")
continue
scan_id = row.get("scan_id")
if not scan_id or not isinstance(scan_id, str):
errors.append(f"{prefix} missing valid string 'scan_id'.")
continue
if scan_id not in gt_cases:
errors.append(f"{prefix} scan_id '{scan_id}' not found in held-out benchmark case list.")
continue
finding_id = row.get("finding_id")
if not finding_id or not isinstance(finding_id, str):
errors.append(f"{prefix} missing valid string 'finding_id'.")
continue
if finding_id in seen_finding_ids:
errors.append(f"{prefix} duplicate finding_id '{finding_id}' in submission.")
continue
seen_finding_ids.add(finding_id)
cls_name = row.get("class")
if cls_name not in ("2c", "2d"):
errors.append(
f"{prefix} (finding_id '{finding_id}') class must be '2c' or '2d', got: {repr(cls_name)}"
)
continue
mask_path_str = row.get("mask_path")
if not mask_path_str or not isinstance(mask_path_str, str):
errors.append(f"{prefix} (finding_id '{finding_id}') missing string 'mask_path'.")
continue
is_soft_mask = bool(row.get("is_soft_mask", False))
if "is_soft_mask" in row and not isinstance(row["is_soft_mask"], bool):
errors.append(f"{prefix} (finding_id '{finding_id}') 'is_soft_mask' must be boolean if specified.")
mask_file = Path(mask_path_str)
if files_dir is not None and not mask_file.is_absolute():
mask_file = files_dir / mask_file
if not mask_file.exists():
errors.append(
f"{prefix} (finding_id '{finding_id}') mask file not found at: '{mask_path_str}'"
)
continue
mask_data, mask_shape, mask_spacing, mask_axcodes, load_err = _load_nifti_data_and_meta(mask_file)
if load_err:
errors.append(
f"{prefix} (finding_id '{finding_id}') failed to load NIfTI mask: {load_err}"
)
continue
gt_meta = gt_cases[scan_id]
expected_shape = tuple(gt_meta["shape"])
expected_spacing = tuple(float(s) for s in gt_meta["spacing"])
expected_axcodes = tuple(gt_meta.get("orientation", ("R", "A", "S")))
if mask_shape != expected_shape:
errors.append(
f"{prefix} (finding_id '{finding_id}') mask shape {mask_shape} does not match expected {expected_shape} for scan '{scan_id}'."
)
continue
spacing_mismatch = False
if len(mask_spacing) != len(expected_spacing):
spacing_mismatch = True
else:
for s1, s2 in zip(mask_spacing, expected_spacing):
if abs(s1 - s2) > 1e-3:
spacing_mismatch = True
break
if spacing_mismatch:
errors.append(
f"{prefix} (finding_id '{finding_id}') mask spacing {mask_spacing} does not match expected {expected_spacing} for scan '{scan_id}' (tolerance 1e-3 mm)."
)
continue
if mask_axcodes and expected_axcodes and mask_axcodes != expected_axcodes:
got_str = _format_axcodes(mask_axcodes)
exp_str = _format_axcodes(expected_axcodes)
errors.append(
f"{prefix} (finding_id '{finding_id}') mask orientation '{got_str}' does not match canonical '{exp_str}' for scan '{scan_id}'."
)
continue
if not is_soft_mask:
is_binary = bool(np.all((mask_data == 0) | (mask_data == 1)))
if not is_binary:
errors.append(
f"{prefix} (finding_id '{finding_id}') declared hard mask (is_soft_mask=false) but contains non-binary values."
)
continue
else:
min_val = float(np.min(mask_data))
max_val = float(np.max(mask_data))
if min_val < -1e-5 or max_val > 1.0 + 1e-5:
errors.append(
f"{prefix} (finding_id '{finding_id}') soft mask values out of range [0, 1] (min={min_val:.3f}, max={max_val:.3f})."
)
continue
has_mask_pos = bool(np.any(mask_data > 0))
del mask_data # free the full-res array immediately (205+ MB per mask)
attr_file = None
attr_path_str = row.get("attribution_map_path")
if attr_path_str:
attr_file = Path(attr_path_str)
if files_dir is not None and not attr_file.is_absolute():
attr_file = files_dir / attr_file
if not attr_file.exists():
errors.append(
f"{prefix} (finding_id '{finding_id}') attribution map file not found at: '{attr_path_str}'"
)
continue
attr_data, a_shape, a_spacing, a_axcodes, a_load_err = _load_nifti_data_and_meta(attr_file)
if a_load_err:
errors.append(
f"{prefix} (finding_id '{finding_id}') failed to load attribution map NIfTI: {a_load_err}"
)
continue
if a_shape != expected_shape:
errors.append(
f"{prefix} (finding_id '{finding_id}') attribution map shape {a_shape} does not match expected {expected_shape}."
)
continue
gt_mask_path = None
gt_dir_base = Path(gt_index.get("gt_dir", "gt_cache")) if isinstance(gt_index, dict) else Path("gt_cache")
for fid, fmeta in gt_meta.get("findings", {}).items():
if fmeta.get("class") == cls_name and fmeta.get("mask_path"):
candidate = gt_dir_base / fmeta["mask_path"]
if candidate.exists():
gt_mask_path = candidate
break
gt_mask_direct = gt_meta.get("gt_masks", {}).get(cls_name) if isinstance(gt_meta.get("gt_masks"), dict) else None
morphology = "focal" if cls_name == "2d" else "non_focal"
loc_sample = LocalizationSample(
case_id=scan_id,
finding_id=finding_id,
model_name=model_name,
class_name=cls_name,
pred_mask=None,
gt_mask=gt_mask_direct,
pred_mask_path=mask_file,
gt_mask_path=gt_mask_path,
gt_shape=expected_shape,
spacing=expected_spacing,
is_soft_mask=is_soft_mask,
morphology=morphology,
dataset="cad_benchmark",
)
parsed_loc_items.append(loc_sample)
if attr_file is not None:
y_true_cls = int(gt_meta.get("labels", {}).get(cls_name, 0))
if scan_id in cls_predictions_by_scan:
y_pred_cls = cls_predictions_by_scan[scan_id][cls_name]
else:
y_pred_cls = 1 if has_mask_pos else 0
exp_sample = ExplainabilitySample(
case_id=scan_id,
finding_id=finding_id,
model_name=model_name,
class_name=cls_name,
attribution_map=None,
gt_mask=gt_mask_direct,
attr_map_path=attr_file,
gt_mask_path=gt_mask_path,
gt_shape=expected_shape,
y_true=y_true_cls,
y_pred=y_pred_cls,
)
parsed_exp_items.append(exp_sample)
if errors:
return ValidationResult(is_valid=False, errors=errors, model_name=model_name, model_type=model_type)
parsed_cls_items = []
if classification_rows:
for scan_id, scores in cls_predictions_by_scan.items():
gt_labels = gt_cases[scan_id]["labels"]
cls_sample = ClassificationSample(
case_id=scan_id,
model_name=model_name,
y_true={"2c": int(gt_labels.get("2c", 0)), "2d": int(gt_labels.get("2d", 0))},
y_score={"2c": int(scores["2c"]), "2d": int(scores["2d"])},
dataset="cad_benchmark",
score_type="hard_label",
)
parsed_cls_items.append(cls_sample)
summary = {
"model_name": model_name,
"model_type": model_type,
"n_classification_cases": len(parsed_cls_items),
"n_segmentation_findings": len(parsed_loc_items),
"n_explainability_findings": len(parsed_exp_items),
}
return ValidationResult(
is_valid=True,
errors=[],
model_name=model_name,
model_type=model_type,
classification_samples=parsed_cls_items if parsed_cls_items else None,
localization_samples=parsed_loc_items if parsed_loc_items else None,
explainability_samples=parsed_exp_items if parsed_exp_items else None,
summary=summary,
)