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)