Spaces:
Running on Zero
Running on Zero
Download validator.py from Gotech/cad-benchmark: direct link, hf CLI and curl.
- Browser
- Download file 16.4 kB
-
https://huggingface.co/spaces/Gotech/cad-benchmark/resolve/main/validator.py
- Command line
-
hf download hf://spaces/Gotech/cad-benchmark/validator.py
-
curl -L -o validator.py https://huggingface.co/spaces/Gotech/cad-benchmark/resolve/main/validator.py
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 | |
| 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, | |
| ) | |