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, )