Source code for zrad.batch.filtering

import logging
from dataclasses import dataclass
from numbers import Integral
from pathlib import Path
from typing import Callable, Sequence

from joblib import Parallel, delayed

from ..exceptions import DataStructureError, InvalidInputParametersError
from ..filtering import create_filter
from ..image import Image
from ._utils import (
    find_nifti_file,
    joblib_parallel_kwargs,
    joblib_progress,
    normalize_common_batch_options,
    normalize_optional_text,
    require_text,
    resolve_patient_folders,
)
from .results import BatchResult

logger = logging.getLogger(__name__)


[docs] @dataclass class FilteringCaseResult: """Per-case result returned by ``BatchFilter``. Attributes ---------- case_name : str Name of the case folder. status : {"processed", "skipped", "failed"} Filtering status for the case. output_path : pathlib.Path or None, optional Path to the written filtered image when filtering succeeds. error : str or None, optional Case-level error message when image loading or filtering fails. """ case_name: str status: str output_path: Path | None = None error: str | None = None
[docs] @dataclass class BatchFilter: """Run one image filter over many case folders and write NIfTI outputs. ``BatchFilter`` is the batch counterpart to the lower-level filtering classes. It discovers case folders, loads one image per case, applies the configured filter, and writes one filtered NIfTI image per processed case. The API is save-to-disk only. Parameters ---------- input_directory : str or pathlib.Path Directory containing one subfolder per case. output_directory : str or pathlib.Path Directory where filtered case folders are written. input_data_type : {"dicom", "nifti"} Input format. Values are normalized to lower-case during validation. modality : {"CT", "MRI", "PET", "MG", "US", "RTDOSE"} Image modality used by the image reader. filter_type : {"Mean", "Laplacian of Gaussian", "Riesz-transformed LoG", "Laws Kernels", "Gabor", "Wavelets", "Simoncelli"} Filter family to apply. filter_dimension : {"2D", "3D"} Apply the filter slice-wise in 2D or volumetrically in 3D. padding_type : str Boundary handling mode passed to the selected filter. number_of_threads : int, optional Number of cases to process in parallel. The default is ``1``. patient_folders : sequence of str or str, optional Explicit case folders to process. Comma-separated strings are accepted. start_folder, stop_folder : str or int, optional Inclusive numeric folder range. Both values must be provided together. nifti_image_name : str, optional Image file name or stem used for NIfTI input. mean_support : int or str, optional Mean-filter kernel side length in voxels. Required for Mean filtering. log_sigma : float or str, optional Gaussian standard deviation in millimetres. Required for LoG and Riesz-transformed LoG filtering. log_cutoff : float or str, optional LoG kernel truncation radius in multiples of ``log_sigma``. laws_response_map : str, optional Laws kernel combination, such as ``"L5E5"`` in 2D or ``"L5E5S5"`` in 3D. laws_rotation_invariance : bool or str, default=False Combine Laws responses over axis permutations and flips. GUI-style ``"Enable"`` and ``"Disable"`` values are also accepted. laws_pooling : {"avg", "max"}, optional Pooling rule for rotation-invariant Laws responses. Required for Laws batch configuration, including when rotation invariance is disabled. laws_energy_map : bool or str, default=False Return a local mean absolute Laws response. Accepts booleans or GUI-style ``"Enable"`` and ``"Disable"`` values. laws_distance : int or str, optional Energy-map neighbourhood radius in voxels. A positive value is required for Laws batch configuration, including when energy maps are disabled. wavelet_response_map : str, optional Separable-wavelet low/high-pass combination, such as ``"LH"`` in 2D or ``"LLH"`` in 3D. wavelet_type : {"db2", "db3", "coif1", "haar"}, optional Wavelet family for separable filtering. wavelet_decomposition_level : int or str, optional Scale level, starting at 1. Required for both separable wavelets (levels 1 or 2) and Simoncelli filtering (any supported positive level). wavelet_rotation_invariance : bool or str, default=False Average separable-wavelet responses over rotations. Accepts booleans or GUI-style ``"Enable"`` and ``"Disable"`` values. gabor_res_mm : float or str, optional Gabor voxel spacing in millimetres per pixel, used to convert physical scales to kernel coordinates. gabor_sigma_mm : float or str, optional Gabor Gaussian envelope standard deviation in millimetres. gabor_lambda_mm : float or str, optional Gabor sinusoidal wavelength in millimetres. gabor_gamma : float or str, optional Gabor kernel aspect ratio. gabor_theta : float or str, optional Gabor orientation angle in radians, or angular step when rotation invariance is enabled. All five numeric Gabor parameters are required when selecting Gabor filtering. gabor_rotation_invariance : bool or str, default=False Average Gabor responses over orientations. Accepts booleans or GUI-style ``"Enable"`` and ``"Disable"`` values. gabor_orthogonal_planes : bool or str, default=False Average Gabor responses across the three orthogonal slice planes. Accepts booleans or GUI-style ``"Enable"`` and ``"Disable"`` values. riesz_order : sequence of int or str, optional Non-negative Riesz multi-index in physical ``(x, y)`` or ``(x, y, z)`` order. Required with positive total order for Riesz-transformed LoG; optional for Simoncelli, where omission or all zeros selects the isotropic response. Comma-separated strings are accepted. structure_tensor_sigma_mm : float or str, optional Positive scale in millimetres for local alignment of pure second-order 3D Riesz-transformed LoG responses, such as ``(2, 0, 0)``. parallel_backend : {"processes", "threads"}, optional Joblib backend preference used when ``number_of_threads`` is greater than one. The default is ``"processes"``. Notes ----- ``validate()`` normalizes public attributes in place. After validation, directories are ``Path`` objects, ``input_data_type`` is lower-case, modality is upper-case, and numeric filter settings are converted to numeric Python values. """ input_directory: str | Path output_directory: str | Path input_data_type: str modality: str filter_type: str filter_dimension: str padding_type: str number_of_threads: int = 1 patient_folders: Sequence[str] | None = None start_folder: str | int | None = None stop_folder: str | int | None = None nifti_image_name: str | None = None mean_support: int | str | None = None log_sigma: float | str | None = None log_cutoff: float | str | None = None laws_response_map: str | None = None laws_rotation_invariance: bool | str = False laws_pooling: str | None = None laws_energy_map: bool | str = False laws_distance: int | str | None = None wavelet_response_map: str | None = None wavelet_type: str | None = None wavelet_decomposition_level: int | str | None = None wavelet_rotation_invariance: bool | str = False gabor_res_mm: float | str | None = None gabor_sigma_mm: float | str | None = None gabor_lambda_mm: float | str | None = None gabor_gamma: float | str | None = None gabor_theta: float | str | None = None gabor_rotation_invariance: bool | str = False gabor_orthogonal_planes: bool | str = False riesz_order: Sequence[int] | str | None = None structure_tensor_sigma_mm: float | str | None = None parallel_backend: str = 'processes'
[docs] def validate(self) -> None: """Validate and normalize filtering configuration. Raises ------ InvalidInputParametersError If the input directory, data type, modality, folder selection, threading, backend, filter type, or filter-specific settings are invalid. """ normalize_common_batch_options(self) self.filter_type = require_text(self.filter_type, "filter_type is required.") self.filter_dimension = require_text(self.filter_dimension, "filter_dimension is required.").upper() self.padding_type = require_text(self.padding_type, "padding_type is required.") if self.filter_dimension not in ['2D', '3D']: raise InvalidInputParametersError("filter_dimension must be '2D' or '3D'.") if self.input_data_type == 'nifti': self.nifti_image_name = normalize_optional_text(self.nifti_image_name) if not self.nifti_image_name: raise InvalidInputParametersError("nifti_image_name is required for NIfTI filtering.") self._validate_filter_parameters()
[docs] def plan(self) -> list[str]: """Return the case folders selected for filtering. Returns ------- folders : list of str Deterministically ordered case folder names selected by ``patient_folders`` or the numeric ``start_folder`` / ``stop_folder`` range. If neither option is set, all non-hidden subfolders are returned. Raises ------ InvalidInputParametersError If validation fails before folder selection. """ self.validate() return self._resolve_patient_folders()
[docs] def run(self, progress_callback: Callable[[int], None] | None = None) -> BatchResult: """Run filtering and write filtered NIfTI images. Parameters ---------- progress_callback : callable, optional Function called as ``progress_callback(step_count)`` after cases complete. ``step_count`` may be greater than one during parallel execution. Returns ------- result : BatchResult Aggregate result with one ``FilteringCaseResult`` per selected case. Notes ----- Case-level failures are recorded in the returned result and do not stop the batch. """ self.validate() self.output_directory.mkdir(parents=True, exist_ok=True) patient_folders = self._resolve_patient_folders() if self.number_of_threads == 1: case_results = [] for patient_folder in patient_folders: case_results.append(self._process_case(patient_folder)) if progress_callback: progress_callback(1) else: with joblib_progress(progress_callback): case_results = Parallel(**joblib_parallel_kwargs(self.number_of_threads, self.parallel_backend))( delayed(self._process_case)(patient_folder) for patient_folder in patient_folders ) return BatchResult(workflow='filtering', case_results=case_results)
[docs] def get_output_filename(self) -> str: """Return the GUI-compatible output file name for the configured filter. Returns ------- filename : str File name ending in ``.nii.gz``. Raises ------ InvalidInputParametersError If ``filter_type`` is not supported. """ values = self._filename_values() filter_formats = { 'Mean': "{filter_type}_{filter_dimension}_{filter_mean_support}support_{filter_padding_type}", 'Laplacian of Gaussian': "{filter_type}_{filter_dimension}_{filter_log_sigma}sigma_" "{filter_log_cutoff}cutoff_{filter_padding_type}", 'Laws Kernels': "{filter_type}_{filter_dimension}_{filter_laws_response_map}_" "{filter_laws_rot_inv}_{filter_laws_pooling}_" "{filter_laws_energy_map}_{filter_laws_distance}_{filter_padding_type}", 'Gabor': "Gabor_{filter_dimension}_" "{filter_gabor_res_mm}resmm_" "{filter_gabor_sigma_mm}sigmm_" "{filter_gabor_lambda_mm}lambmm_" "g{filter_gabor_gamma}_" "t{filter_gabor_theta}_" "{filter_gabor_rotinv}_" "{filter_gabor_ortho}_" "{filter_padding_type}", 'Riesz-transformed LoG': "RieszLoG_{filter_dimension}_{filter_log_sigma}sigma_" "{filter_log_cutoff}cutoff_Riesz{filter_riesz_order}" "{filter_structure_tensor_suffix}_{filter_padding_type}", 'Simoncelli': "Simoncelli_{filter_dimension}_{filter_wavelet_decomp_lvl}_" "Riesz{filter_riesz_order}_{filter_padding_type}", } if self.filter_type == 'Wavelets': filename = ( "{filter_wavelet_type}_{filter_dimension}_" "{filter_wavelet_resp_map}_" "{filter_wavelet_decomp_lvl}_" "{filter_wavelet_rot_inv}_" "{filter_padding_type}" ).format(**values) elif self.filter_type in filter_formats: filename = filter_formats[self.filter_type].format(**values) else: raise InvalidInputParametersError(f"Unknown filter type: {self.filter_type}") return f"{filename}.nii.gz"
def _resolve_patient_folders(self) -> list[str]: return resolve_patient_folders( self.input_directory, self.patient_folders, self.start_folder, self.stop_folder, ) def _process_case(self, case_name: str) -> FilteringCaseResult: case_dir = self.input_directory / case_name logger.info("Filtering patient's %s image.", case_name) try: image = self._load_image(case_dir) except (DataStructureError, FileNotFoundError, ValueError) as exc: return FilteringCaseResult(case_name=case_name, status='skipped', error=str(exc)) except Exception as exc: logger.exception("Patient %s failed while loading image.", case_name) return FilteringCaseResult(case_name=case_name, status='failed', error=str(exc)) try: filtering = self._create_filter() image_new = filtering.apply(image) output_path = self.output_directory / case_name / self.get_output_filename() image_new.save_as_nifti(output_path) return FilteringCaseResult(case_name=case_name, status='processed', output_path=output_path) except Exception as exc: logger.exception("Patient %s failed during filtering.", case_name) return FilteringCaseResult(case_name=case_name, status='failed', error=str(exc)) def _load_image(self, case_dir: Path) -> Image: if self.input_data_type == 'dicom': return Image.from_dicom(case_dir, modality=self.modality) image_path = find_nifti_file(case_dir, self.nifti_image_name) if image_path is None: raise FileNotFoundError(case_dir / self.nifti_image_name) return Image.from_nifti(image_path) def _create_filter(self): if self.filter_type == 'Mean': return create_filter( filtering_method='Mean', padding_type=self.padding_type, support=self.mean_support, dimensionality=self.filter_dimension, ) if self.filter_type == 'Laplacian of Gaussian': return create_filter( filtering_method='Laplacian of Gaussian', padding_type=self.padding_type, sigma_mm=self.log_sigma, cutoff=self.log_cutoff, dimensionality=self.filter_dimension, ) if self.filter_type == 'Riesz-transformed LoG': return create_filter( filtering_method=self.filter_type, padding_type=self.padding_type, sigma_mm=self.log_sigma, cutoff=self.log_cutoff, dimensionality=self.filter_dimension, riesz_order=self.riesz_order, structure_tensor_sigma_mm=self.structure_tensor_sigma_mm, ) if self.filter_type == 'Laws Kernels': return create_filter( filtering_method='Laws Kernels', response_map=self.laws_response_map, padding_type=self.padding_type, dimensionality=self.filter_dimension, rotation_invariance=_as_bool(self.laws_rotation_invariance), pooling=self.laws_pooling, energy_map=_as_bool(self.laws_energy_map), distance=self.laws_distance, ) if self.filter_type == 'Gabor': return create_filter( filtering_method='Gabor', padding_type=self.padding_type, res_mm=self.gabor_res_mm, sigma_mm=self.gabor_sigma_mm, lambda_mm=self.gabor_lambda_mm, gamma=self.gabor_gamma, theta=self.gabor_theta, rotation_invariance=_as_bool(self.gabor_rotation_invariance), orthogonal_planes=_as_bool(self.gabor_orthogonal_planes), ) if self.filter_type == 'Wavelets': return create_filter( filtering_method='Wavelets', dimensionality=self.filter_dimension, padding_type=self.padding_type, wavelet_type=self.wavelet_type, response_map=self.wavelet_response_map, decomposition_level=self.wavelet_decomposition_level, rotation_invariance=_as_bool(self.wavelet_rotation_invariance), ) if self.filter_type == 'Simoncelli': return create_filter( filtering_method=self.filter_type, padding_type=self.padding_type, decomposition_level=self.wavelet_decomposition_level, dimensionality=self.filter_dimension, riesz_order=self.riesz_order, ) raise InvalidInputParametersError(f"Filter_type {self.filter_type} not supported.") def _validate_filter_parameters(self) -> None: if self.filter_type == 'Mean': self.mean_support = _require_positive_int(self.mean_support, "mean_support is required.") elif self.filter_type == 'Laplacian of Gaussian': self.log_sigma = _require_float(self.log_sigma, "log_sigma is required.") self.log_cutoff = _require_float(self.log_cutoff, "log_cutoff is required.") elif self.filter_type == 'Riesz-transformed LoG': self.log_sigma = _require_float(self.log_sigma, "log_sigma is required.") self.log_cutoff = _require_float(self.log_cutoff, "log_cutoff is required.") self.riesz_order = _require_riesz_order(self.riesz_order, self.filter_dimension) if self.structure_tensor_sigma_mm is not None and str(self.structure_tensor_sigma_mm).strip(): self.structure_tensor_sigma_mm = _require_float( self.structure_tensor_sigma_mm, "structure_tensor_sigma_mm must be a number." ) if self.structure_tensor_sigma_mm <= 0: raise InvalidInputParametersError('structure_tensor_sigma_mm must be a positive number.') if self.filter_dimension != '3D' or sum(self.riesz_order) != 2: raise InvalidInputParametersError( 'Structure-tensor alignment is supported for second-order 3D Riesz transforms only.' ) if 2 not in self.riesz_order: raise InvalidInputParametersError( 'Structure-tensor alignment supports pure second-order Riesz indices only; ' 'mixed-order indices have sign-ambiguous eigenvector steering.' ) else: self.structure_tensor_sigma_mm = None elif self.filter_type == 'Laws Kernels': self.laws_response_map = require_text(self.laws_response_map, "laws_response_map is required.") self.laws_pooling = require_text(self.laws_pooling, "laws_pooling is required.") self.laws_distance = _require_positive_int(self.laws_distance, "laws_distance is required.") self.laws_rotation_invariance = _normalize_enable_disable(self.laws_rotation_invariance) self.laws_energy_map = _normalize_enable_disable(self.laws_energy_map) elif self.filter_type == 'Gabor': self.gabor_res_mm = _require_float(self.gabor_res_mm, "gabor_res_mm is required.") self.gabor_sigma_mm = _require_float(self.gabor_sigma_mm, "gabor_sigma_mm is required.") self.gabor_lambda_mm = _require_float(self.gabor_lambda_mm, "gabor_lambda_mm is required.") self.gabor_gamma = _require_float(self.gabor_gamma, "gabor_gamma is required.") self.gabor_theta = _require_float(self.gabor_theta, "gabor_theta is required.") self.gabor_rotation_invariance = _normalize_enable_disable(self.gabor_rotation_invariance) self.gabor_orthogonal_planes = _normalize_enable_disable(self.gabor_orthogonal_planes) elif self.filter_type == 'Wavelets': self.wavelet_type = require_text(self.wavelet_type, "wavelet_type is required.") self.wavelet_response_map = require_text(self.wavelet_response_map, "wavelet_response_map is required.") self.wavelet_decomposition_level = _require_positive_int( self.wavelet_decomposition_level, "wavelet_decomposition_level is required.", ) self.wavelet_rotation_invariance = _normalize_enable_disable(self.wavelet_rotation_invariance) elif self.filter_type == 'Simoncelli': if self.padding_type not in ('nearest', 'wrap', 'periodic'): raise InvalidInputParametersError("Simoncelli padding_type must be 'nearest', 'wrap', or 'periodic'.") self.wavelet_decomposition_level = _require_positive_int( self.wavelet_decomposition_level, "wavelet_decomposition_level is required." ) self.riesz_order = _optional_riesz_order(self.riesz_order, self.filter_dimension) else: raise InvalidInputParametersError(f"Filter_type {self.filter_type} not supported.") def _filename_values(self) -> dict: return { 'filter_type': self.filter_type, 'filter_dimension': self.filter_dimension, 'filter_padding_type': self.padding_type, 'filter_mean_support': self.mean_support, 'filter_log_sigma': self.log_sigma, 'filter_log_cutoff': self.log_cutoff, 'filter_laws_response_map': self.laws_response_map, 'filter_laws_rot_inv': _enable_disable_text(self.laws_rotation_invariance), 'filter_laws_pooling': self.laws_pooling, 'filter_laws_energy_map': _enable_disable_text(self.laws_energy_map), 'filter_laws_distance': self.laws_distance, 'filter_wavelet_type': self.wavelet_type, 'filter_wavelet_resp_map': self.wavelet_response_map, 'filter_wavelet_decomp_lvl': self.wavelet_decomposition_level, 'filter_wavelet_rot_inv': _enable_disable_text(self.wavelet_rotation_invariance), 'filter_gabor_res_mm': self.gabor_res_mm, 'filter_gabor_sigma_mm': self.gabor_sigma_mm, 'filter_gabor_lambda_mm': self.gabor_lambda_mm, 'filter_gabor_gamma': self.gabor_gamma, 'filter_gabor_theta': self.gabor_theta, 'filter_gabor_rotinv': _enable_disable_text(self.gabor_rotation_invariance), 'filter_gabor_ortho': _enable_disable_text(self.gabor_orthogonal_planes), 'filter_riesz_order': 'none' if self.riesz_order is None else '-'.join(map(str, self.riesz_order)), 'filter_structure_tensor_suffix': ( '' if self.structure_tensor_sigma_mm is None else f'_Tensor{self.structure_tensor_sigma_mm}sigma' ), }
def _require_positive_int(value, message: str) -> int: if value is None or str(value).strip() == '': raise InvalidInputParametersError(message) try: result = int(value) except (TypeError, ValueError): raise InvalidInputParametersError(message) if result <= 0: raise InvalidInputParametersError(message) return result def _require_float(value, message: str) -> float: if value is None or str(value).strip() == '': raise InvalidInputParametersError(message) try: return float(value) except (TypeError, ValueError): raise InvalidInputParametersError(message) def _optional_riesz_order(value, dimensionality: str): if value is None or str(value).strip() == '': return None return _require_riesz_order(value, dimensionality, allow_zero=True) def _require_riesz_order(value, dimensionality: str, *, allow_zero: bool = False) -> tuple[int, ...]: message = f"riesz_order must contain {dimensionality[0]} comma-separated non-negative integers." try: if isinstance(value, str): values = [item.strip() for item in value.split(',')] else: values = tuple(value) if any(not isinstance(item, Integral) or isinstance(item, bool) for item in values): raise InvalidInputParametersError(message) result = tuple(int(item) for item in values) except (TypeError, ValueError): raise InvalidInputParametersError(message) if len(result) != int(dimensionality[0]) or any(item < 0 for item in result): raise InvalidInputParametersError(message) if not allow_zero and sum(result) == 0: raise InvalidInputParametersError("riesz_order must have a positive total order.") return result def _normalize_enable_disable(value) -> str: if isinstance(value, str): text = value.strip() if text in ['Enable', 'Disable']: return text if text.lower() in ['true', 'yes', '1']: return 'Enable' if text.lower() in ['false', 'no', '0']: return 'Disable' return 'Enable' if bool(value) else 'Disable' def _as_bool(value) -> bool: return _normalize_enable_disable(value) == 'Enable' def _enable_disable_text(value) -> str: return _normalize_enable_disable(value)