From ad51b1c6a07cfac2168136031783f15b6bb0e09b Mon Sep 17 00:00:00 2001 From: luandalmazo Date: Thu, 6 Aug 2026 14:52:14 -0300 Subject: [PATCH] add ruff and applied minor fixes --- .github/workflows/lint.yaml | 29 +++ datamint/__init__.py | 17 +- datamint/api/base_api.py | 53 +++--- datamint/api/client.py | 32 ++-- datamint/api/dto/__init__.py | 3 +- datamint/api/dto/annotation_dto.py | 5 +- datamint/api/endpoints/__init__.py | 18 +- datamint/api/endpoints/annotations_api.py | 35 ++-- datamint/api/endpoints/annotationsets_api.py | 13 +- datamint/api/endpoints/channels_api.py | 5 +- datamint/api/endpoints/datasetsinfo_api.py | 8 +- datamint/api/endpoints/deploy_model_api.py | 9 +- datamint/api/endpoints/inference_api.py | 11 +- datamint/api/endpoints/models_api.py | 8 +- datamint/api/endpoints/projects_api.py | 27 ++- datamint/api/endpoints/resources_api.py | 73 +++---- datamint/api/endpoints/users_api.py | 22 ++- datamint/api/entity_base_api.py | 18 +- datamint/client_cmd_tools/datamint_config.py | 19 +- datamint/client_cmd_tools/datamint_example.py | 10 +- .../client_cmd_tools/datamint_inference.py | 24 ++- datamint/client_cmd_tools/datamint_init.py | 3 +- datamint/client_cmd_tools/datamint_train.py | 33 ++-- datamint/client_cmd_tools/datamint_upload.py | 60 +++--- datamint/configs.py | 16 +- datamint/dataset/__init__.py | 21 +-- datamint/dataset/annotation_processor.py | 15 +- datamint/dataset/base.py | 48 ++--- datamint/dataset/factory.py | 10 +- datamint/dataset/image_dataset.py | 6 +- datamint/dataset/multiframe_dataset.py | 8 +- datamint/dataset/sliced_dataset.py | 41 ++-- datamint/dataset/sliced_video_dataset.py | 32 ++-- datamint/dataset/split_result.py | 2 +- datamint/entities/__init__.py | 25 ++- datamint/entities/annotation_worklist.py | 2 +- datamint/entities/annotations/__init__.py | 14 +- datamint/entities/annotations/annotation.py | 9 +- .../entities/annotations/annotation_spec.py | 1 + .../entities/annotations/base_segmentation.py | 7 +- .../entities/annotations/box_annotation.py | 5 +- datamint/entities/annotations/geometry.py | 16 +- .../annotations/image_segmentation.py | 2 +- .../entities/annotations/line_annotation.py | 6 +- .../annotations/numeric_annotation.py | 2 +- .../annotations/volume_segmentation.py | 13 +- datamint/entities/base_entity.py | 17 +- datamint/entities/cache_manager.py | 16 +- datamint/entities/channel.py | 5 +- datamint/entities/datasetinfo.py | 8 +- datamint/entities/deployjob.py | 3 +- datamint/entities/inferencejob.py | 8 +- datamint/entities/project.py | 15 +- datamint/entities/resource.py | 48 +++-- datamint/entities/resources/nifti_resource.py | 3 +- datamint/entities/resources/video_resource.py | 4 +- .../entities/resources/volume_resource.py | 11 +- datamint/entities/sliced_resource.py | 10 +- datamint/entities/sliced_resource_base.py | 4 +- datamint/entities/sliced_video_resource.py | 7 +- datamint/examples/__init__.py | 7 +- datamint/examples/example_projects.py | 10 +- datamint/examples/synapse_dataset.py | 2 +- datamint/exceptions.py | 6 - datamint/lightning/__init__.py | 30 +-- datamint/lightning/datamodule.py | 4 +- datamint/lightning/trainers/__init__.py | 34 ++-- datamint/lightning/trainers/base_trainer.py | 71 +++---- .../trainers/classification_trainer.py | 27 +-- .../lightning/trainers/detection_trainer.py | 9 +- .../trainers/lightning_modules/__init__.py | 12 +- .../trainers/lightning_modules/base.py | 15 +- .../classification_module.py | 10 +- .../detection_modules/yolox_module.py | 2 + .../lightning_modules/segmentation_module.py | 15 +- .../segmentation_modules/__init__.py | 6 +- .../segmentation_modules/deeplabv3plus.py | 8 +- .../segmentation_modules/smp_module.py | 3 - .../segmentation_modules/transunet.py | 11 +- .../segmentation_modules/unetpp.py | 12 +- .../segmentation_modules/unetrpp.py | 21 ++- datamint/lightning/trainers/seg2d_trainer.py | 44 ++--- datamint/lightning/trainers/seg3d_trainer.py | 24 ++- .../trainers/segmentation_trainer.py | 3 +- .../trainers/specialized/__init__.py | 13 +- .../trainers/specialized/deeplabv3plus.py | 9 +- .../trainers/specialized/efficientnetv2.py | 5 +- .../nnunet/_nnunet_trainer_bridge.py | 7 +- .../specialized/nnunet/data_export.py | 3 +- .../specialized/nnunet/data_import.py | 1 + .../specialized/nnunet/inference_model.py | 2 +- .../trainers/specialized/nnunet/trainer.py | 27 ++- .../trainers/specialized/transunet.py | 11 +- .../lightning/trainers/specialized/unetpp.py | 13 +- .../lightning/trainers/specialized/unetrpp.py | 12 +- .../lightning/trainers/specialized/yolox.py | 21 ++- .../lightning/trainers/vol_seg_trainer.py | 36 ++-- datamint/mlflow/__init__.py | 16 +- datamint/mlflow/artifact/__init__.py | 2 +- datamint/mlflow/data/__init__.py | 4 +- datamint/mlflow/data/datamint_dataset.py | 6 +- datamint/mlflow/env_utils.py | 8 +- datamint/mlflow/env_vars.py | 1 + datamint/mlflow/flavors/__init__.py | 27 +-- datamint/mlflow/flavors/datamint_flavor.py | 51 ++--- datamint/mlflow/flavors/model.py | 16 +- datamint/mlflow/flavors/prediction_router.py | 7 +- datamint/mlflow/flavors/validation.py | 5 +- .../mlflow/lightning/callbacks/__init__.py | 8 +- .../lightning/callbacks/modelcheckpoint.py | 44 ++--- datamint/mlflow/models/__init__.py | 9 +- .../mlflow/models/datamint_model_store.py | 18 +- datamint/mlflow/store_utils.py | 6 +- datamint/mlflow/tracking/datamint_store.py | 16 +- .../mlflow/tracking/default_experiment.py | 7 +- datamint/mlflow/tracking/fluent.py | 9 +- datamint/types.py | 6 +- datamint/utils/annotation_agreement.py | 3 +- datamint/utils/collection_utils.py | 2 +- datamint/utils/logging_utils.py | 20 +- datamint/utils/nifti_utils.py | 5 +- datamint/utils/torchmetrics.py | 5 +- datamint/utils/visualization.py | 17 +- .../01_getting_started/01_upload_data.ipynb | 33 +++- .../01_getting_started/02_explore_data.ipynb | 42 ++++- .../01_project_scoped_splits.ipynb | 1 + .../03_datasets/02_patient_wise_splits.ipynb | 3 +- notebooks/03_datasets/03_build_dataset.ipynb | 2 +- notebooks/03_datasets/04_volume_dataset.ipynb | 3 +- .../01_mlflow_manual_logging.ipynb | 1 + .../01_deploy_registered_model.ipynb | 14 +- .../02_deploy_external_model.ipynb | 21 ++- .../05_deployment/03_validate_model.ipynb | 25 ++- ...04_predict_images_volumes_and_videos.ipynb | 9 +- .../full_3d/01_synapse_unetrpp.ipynb | 9 +- .../full_3d/02_synapse_nnunet.ipynb | 68 ++++++- .../01_fracatlas_classification.ipynb | 18 +- .../slice_based/02_busi_segmentation.ipynb | 28 ++- .../slice_based/03_bccd_detection.ipynb | 18 +- pyproject.toml | 15 +- tests/conftest.py | 1 - tests/test_annotation_agreement.py | 6 +- tests/test_api_handler.py | 22 ++- tests/test_auto_config_mlflow.py | 8 +- tests/test_datamint_config.py | 5 +- tests/test_dataset_patient_split.py | 1 - tests/test_detection_trainer.py | 3 +- tests/test_imports.py | 3 +- tests/test_nnunet_data_export.py | 29 ++- tests/test_nnunet_data_import.py | 13 +- tests/test_nnunet_inference_model.py | 12 +- tests/test_nnunet_integration.py | 5 +- tests/test_nnunet_trainer.py | 8 +- tests/test_nnunet_trainer_bridge.py | 128 ------------- tests/test_upload_validation.py | 7 +- tests/test_yolox_module.py | 178 ------------------ tests/test_yolox_trainer.py | 8 +- 157 files changed, 1342 insertions(+), 1264 deletions(-) create mode 100644 .github/workflows/lint.yaml delete mode 100644 tests/test_nnunet_trainer_bridge.py delete mode 100644 tests/test_yolox_module.py diff --git a/.github/workflows/lint.yaml b/.github/workflows/lint.yaml new file mode 100644 index 00000000..3b49a82a --- /dev/null +++ b/.github/workflows/lint.yaml @@ -0,0 +1,29 @@ +name: Lint + +on: + push: + branches: [main] + pull_request: + branches: [main, fix/*, hotfix/*, release/*, develop] + +permissions: + contents: read + +jobs: + ruff: + runs-on: ubuntu-latest + steps: + - name: Checkout code + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.12" + + - name: Install ruff + run: pip install "ruff==0.16.1" + + - name: Run ruff check + run: ruff check . + continue-on-error: true diff --git a/datamint/__init__.py b/datamint/__init__.py index e1399afd..c6c45e84 100644 --- a/datamint/__init__.py +++ b/datamint/__init__.py @@ -4,16 +4,23 @@ import importlib.metadata from typing import TYPE_CHECKING + from .utils.logging_utils import setup_file_logging_if_enabled setup_file_logging_if_enabled() if TYPE_CHECKING: - from .api.client import Api + from .api.client import Api as Api + # New modular datasets - from .dataset.image_dataset import ImageDataset - from .dataset.volume_dataset import VolumeDataset - from .mlflow.flavors.validation import validate_model, ValidationReport, ValidationIssue, ModelValidationError - from .default_project import select_project + from .dataset.image_dataset import ImageDataset as ImageDataset + from .dataset.volume_dataset import VolumeDataset as VolumeDataset + from .default_project import select_project as select_project + from .mlflow.flavors.validation import ( + ModelValidationError as ModelValidationError, + ValidationIssue as ValidationIssue, + ValidationReport as ValidationReport, + validate_model as validate_model, + ) else: import lazy_loader as lazy diff --git a/datamint/api/base_api.py b/datamint/api/base_api.py index 636f6666..7c79bc3e 100644 --- a/datamint/api/base_api.py +++ b/datamint/api/base_api.py @@ -1,26 +1,28 @@ +import asyncio +import contextlib +import gzip +import json import logging +import os +from collections.abc import AsyncGenerator, Generator +from dataclasses import dataclass +from io import BytesIO from typing import TYPE_CHECKING -from collections.abc import Generator, AsyncGenerator + +import aiohttp +import cv2 import httpx -from dataclasses import dataclass +from PIL import Image + from datamint.exceptions import ( - ItemNotFoundError, AuthenticationError, - PermissionDeniedError, - ValidationError, + ItemNotFoundError, NetworkError, + PermissionDeniedError, ServerError, + ValidationError, ) -import aiohttp -import json -from PIL import Image -import cv2 -from io import BytesIO -import gzip -import contextlib -import asyncio from datamint.utils.env import ensure_asyncio_loop -import os if TYPE_CHECKING: from datamint.api.client import Api @@ -86,9 +88,10 @@ def __init__(self, self._pid = os.getpid() # Track PID to detect DataLoader worker forks self.client = client or BaseApi._create_client(config) self.semaphore = asyncio.Semaphore(_ASYNC_REQUEST_LIMIT) - self._api_instance: 'Api | None' = None # Injected by Api class + self._api_instance: Api | None = None # Injected by Api class self._aiohttp_connector: aiohttp.TCPConnector | None = None self._aiohttp_session: aiohttp.ClientSession | None = None + self._pending_close_task: asyncio.Task | None = None ensure_asyncio_loop() @staticmethod @@ -160,6 +163,7 @@ def _create_aiohttp_connector(self, force_close: bool = False) -> aiohttp.TCPCon Configured TCPConnector for aiohttp sessions. """ import ssl + import certifi limit = _ASYNC_REQUEST_LIMIT @@ -215,10 +219,10 @@ def _close_aiohttp_session(self) -> None: # If we're in an environment where the loop is running and not patched, # fall back to scheduling the close. try: - loop.create_task(self._aiohttp_session.close()) + self._pending_close_task = loop.create_task(self._aiohttp_session.close()) + self._pending_close_task.add_done_callback(lambda _: setattr(self, '_pending_close_task', None)) except Exception as e: logger.info(f"Unable to schedule aiohttp session close: {e}") - pass finally: self._aiohttp_session = None self._aiohttp_connector = None @@ -305,6 +309,7 @@ def _ensure_client_fresh(self) -> None: # Invalidate any inherited aiohttp session as well. self._aiohttp_session = None self._aiohttp_connector = None + self._pending_close_task = None def _make_request(self, method: str, endpoint: str, **kwargs) -> httpx.Response: """Make HTTP request with error handling and retries. @@ -699,12 +704,12 @@ def convert_format(bytes_array: bytes, >>> dicom = BaseApi.convert_format(dicom_bytes) """ - import pydicom import nibabel as nib + import pydicom from medimgkit.format_detection import GZIP_MIME_TYPES if mimetype is None: - mimetype, ext = BaseApi._determine_mimetype(bytes_array) + mimetype, _ext = BaseApi._determine_mimetype(bytes_array) if mimetype is None: raise ValueError("Could not determine mimetype from content.") content_io = BytesIO(bytes_array) @@ -725,12 +730,12 @@ def convert_format(bytes_array: bytes, ndata = nib.Nifti1Image.from_stream(content_io) ndata.get_fdata() # force loading before IO is closed return ndata - except Exception as e: + except Exception: if file_path is not None: ndata = nib.load(file_path) ndata.get_fdata() # force loading before IO is closed return ndata - raise e + raise elif mimetype in GZIP_MIME_TYPES: # let's hope it's a .nii.gz with gzip.open(content_io, 'rb') as f: @@ -754,7 +759,11 @@ def _determine_mimetype(content: bytes, Returns: Tuple of (inferred_mimetype, file_extension) """ - from medimgkit.format_detection import DEFAULT_MIME_TYPE, guess_typez, guess_extension + from medimgkit.format_detection import ( + DEFAULT_MIME_TYPE, + guess_extension, + guess_typez, + ) # Determine mimetype from file content mimetype_list, ext = guess_typez(content, use_magic=True) mimetype = mimetype_list[-1] diff --git a/datamint/api/client.py b/datamint/api/client.py index cfd17abd..7f5df99f 100644 --- a/datamint/api/client.py +++ b/datamint/api/client.py @@ -1,15 +1,22 @@ -from typing import Any +import logging +from typing import Any, ClassVar -from .base_api import ApiConfig, BaseApi -from .endpoints import (ProjectsApi, ResourcesApi, AnnotationsApi, - ChannelsApi, UsersApi, DatasetsInfoApi, - AnnotationWorklistApi, DeployModelApi, - InferenceApi - ) -from .endpoints.models_api import ModelsApi import datamint.configs from datamint.exceptions import AuthenticationError, NetworkError -import logging + +from .base_api import ApiConfig, BaseApi +from .endpoints import ( + AnnotationsApi, + AnnotationWorklistApi, + ChannelsApi, + DatasetsInfoApi, + DeployModelApi, + InferenceApi, + ProjectsApi, + ResourcesApi, + UsersApi, +) +from .endpoints.models_api import ModelsApi _LOGGER = logging.getLogger(__name__) @@ -19,7 +26,7 @@ class Api: DEFAULT_SERVER_URL = 'https://api.datamint.io' DATAMINT_API_VENV_NAME = datamint.configs.ENV_VARS[datamint.configs.APIKEY_KEY] - _API_MAP: dict[str, type[BaseApi]] = { + _API_MAP: ClassVar[dict[str, type[BaseApi]]] = { 'projects': ProjectsApi, 'resources': ResourcesApi, 'annotations': AnnotationsApi, @@ -37,7 +44,7 @@ class Api: # (e.g. one per DataLoader worker or dataset auto-refresh) once a given # configuration is known to work. A changed value simply misses the cache, # forcing a fresh check. - _verified_connections: set[tuple[str, str | None, bool | str]] = set() + _verified_connections: ClassVar[set[tuple[str, str | None, bool | str]]] = set() def __init__(self, server_url: str | None = None, @@ -67,7 +74,7 @@ def __init__(self, if api_key is None: api_key = datamint.configs.get_value(datamint.configs.APIKEY_KEY) if api_key is None: - msg = f"API key not provided! Use the environment variable " + \ + msg = "API key not provided! Use the environment variable " + \ f"{Api.DATAMINT_API_VENV_NAME} or pass it as an argument." raise AuthenticationError(msg) self.config = ApiConfig( @@ -121,7 +128,6 @@ def close(self) -> None: endpoint.close() except Exception as e: _LOGGER.warning(f"Error closing endpoint {endpoint}: {e}") - pass # Close shared httpx clients owned by this Api for client in (self._client, self._highclient, self._mlclient): diff --git a/datamint/api/dto/__init__.py b/datamint/api/dto/__init__.py index 29ffd069..5ae550e4 100644 --- a/datamint/api/dto/__init__.py +++ b/datamint/api/dto/__init__.py @@ -2,8 +2,7 @@ CreateAnnotationDto, ) - __all__ = [ - "annotation_dto", "CreateAnnotationDto", + "annotation_dto", ] diff --git a/datamint/api/dto/annotation_dto.py b/datamint/api/dto/annotation_dto.py index d78bf4d5..137f9db7 100644 --- a/datamint/api/dto/annotation_dto.py +++ b/datamint/api/dto/annotation_dto.py @@ -14,7 +14,8 @@ CreateAnnotationDto: Main DTO for creating annotation requests. """ -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any + from datamint.entities.annotations import AnnotationType if TYPE_CHECKING: @@ -50,7 +51,7 @@ def __init__(self, self.units = units self.model_id = model_id if model_id is not None: - if is_model == False: + if is_model is False: raise ValueError("model_id==False while self.model_id is provided.") if not isinstance(model_id, str): raise ValueError("model_id must be a string if provided.") diff --git a/datamint/api/endpoints/__init__.py b/datamint/api/endpoints/__init__.py index a1d69c43..a708a5e6 100644 --- a/datamint/api/endpoints/__init__.py +++ b/datamint/api/endpoints/__init__.py @@ -1,23 +1,23 @@ """API endpoint handlers.""" from .annotations_api import AnnotationsApi +from .annotationsets_api import AnnotationWorklistApi from .channels_api import ChannelsApi -from .projects_api import ProjectsApi -from .resources_api import ResourcesApi -from .users_api import UsersApi from .datasetsinfo_api import DatasetsInfoApi -from .annotationsets_api import AnnotationWorklistApi from .deploy_model_api import DeployModelApi from .inference_api import InferenceApi +from .projects_api import ProjectsApi +from .resources_api import ResourcesApi +from .users_api import UsersApi __all__ = [ + 'AnnotationWorklistApi', 'AnnotationsApi', - 'ChannelsApi', - 'ProjectsApi', - 'ResourcesApi', - 'UsersApi', + 'ChannelsApi', 'DatasetsInfoApi', - 'AnnotationWorklistApi', 'DeployModelApi', 'InferenceApi', + 'ProjectsApi', + 'ResourcesApi', + 'UsersApi', ] diff --git a/datamint/api/endpoints/annotations_api.py b/datamint/api/endpoints/annotations_api.py index f0789b36..48b0380e 100644 --- a/datamint/api/endpoints/annotations_api.py +++ b/datamint/api/endpoints/annotations_api.py @@ -1,20 +1,20 @@ -from typing import Any, BinaryIO, IO, Literal, overload, TYPE_CHECKING -from collections.abc import Generator, Sequence -from datetime import date -from io import BytesIO -from pathlib import Path import asyncio import json import logging import os +from collections.abc import Generator, Sequence +from datetime import date +from io import BytesIO +from pathlib import Path +from typing import IO, TYPE_CHECKING, Any, BinaryIO, Literal, overload import aiohttp import httpx import numpy as np import pydicom +from medimgkit import ViewPlane from medimgkit.format_detection import guess_type from medimgkit.nifti_utils import DEFAULT_NIFTI_MIME -from medimgkit import ViewPlane from nibabel.loadsave import load as nib_load from nibabel.nifti1 import Nifti1Image from PIL import Image @@ -38,10 +38,11 @@ from ..entity_base_api import ApiConfig, CreatableEntityApi, DeletableEntityApi if TYPE_CHECKING: - from .resources_api import ResourcesApi - from .models_api import ModelsApi from datamint.entities import Project + from .models_api import ModelsApi + from .resources_api import ResourcesApi + _LOGGER = logging.getLogger(__name__) _USER_LOGGER = logging.getLogger('user_logger') @@ -250,7 +251,7 @@ async def _upload_segmentations_async(self, List of annotation IDs created. """ if upload_volume == 'auto': - if isinstance(file_path, str) and (file_path.endswith('.nii') or file_path.endswith('.nii.gz')): + if isinstance(file_path, str) and (file_path.endswith(('.nii', '.nii.gz'))): upload_volume = True else: upload_volume = False @@ -392,7 +393,7 @@ async def _upload_single_frame_segmentation_async(self, annotations.append(ann) # Validate unique identifiers - if len(annotations) != len(set([a.identifier for a in annotations])): + if len(annotations) != len({a.identifier for a in annotations}): raise ValueError( "Multiple annotations with the same identifier, frame_index, scope and author is not supported yet." ) @@ -733,7 +734,7 @@ def upload_segmentations(self, if isinstance(file_path, str) and not os.path.exists(file_path): raise FileNotFoundError(f"File {file_path} not found.") - if isinstance(file_path, str) and (file_path.endswith('.nii') or file_path.endswith('.nii.gz')): + if isinstance(file_path, str) and (file_path.endswith(('.nii', '.nii.gz'))): raise ValueError( "NIfTI files are volume segmentations. Use `upload_volume_segmentation` instead." ) @@ -895,7 +896,7 @@ async def _upload_volume_segmentation_async(self, if isinstance(name, str): raise NotImplementedError("`name=string` is not supported yet for volume segmentation.") if isinstance(name, dict): - if any(isinstance(k, tuple) for k in name.keys()): + if any(isinstance(k, tuple) for k in name): raise NotImplementedError( "For volume segmentations, `name` must be a dictionary with integer keys only.") if 'default' in name: @@ -905,7 +906,7 @@ async def _upload_volume_segmentation_async(self, # Prepare file for upload if isinstance(file_path, str): - if file_path.endswith('.nii') or file_path.endswith('.nii.gz'): + if file_path.endswith(('.nii', '.nii.gz')): # Upload NIfTI file directly _LOGGER.debug('uploading segmentation as a volume') with open(file_path, 'rb') as f: @@ -998,7 +999,7 @@ def _generate_segmentations_ios(file_path: str | np.ndarray, nframes = normalized_imgs.shape[3] fios = AnnotationsApi._numpy_to_bytesio_png(normalized_imgs) - elif file_path.endswith('.nii') or file_path.endswith('.nii.gz'): + elif file_path.endswith(('.nii', '.nii.gz')): loaded_image: Any = nib_load(file_path) segs_imgs = loaded_image.get_fdata() if segs_imgs.ndim != 3 and segs_imgs.ndim != 2: @@ -1118,7 +1119,7 @@ def create_image_classification(self, def create_numeric_annotation(self, resource: str | Resource, identifier: str, - value: int | float, + value: float, units: str | None = None, imported_from: str | None = None, model_id: str | None = None, @@ -1497,7 +1498,7 @@ async def _async_download_segmentation_file(self, try: resource_id = self.get_by_id(annotation_id).resource_id except Exception as e: - error_msg = f"Failed to get resource_id for annotation {annotation_id}: {str(e)}" + error_msg = f"Failed to get resource_id for annotation {annotation_id}: {e!s}" _LOGGER.error(error_msg) if progress_bar: progress_bar.update(1) @@ -1514,7 +1515,7 @@ async def _async_download_segmentation_file(self, progress_bar.update(1) return {'success': True, 'annotation_id': annotation_id} except Exception as e: - error_msg = f"Failed to download annotation {annotation_id}: {str(e)}" + error_msg = f"Failed to download annotation {annotation_id}: {e!s}" _LOGGER.error(error_msg) if progress_bar: progress_bar.update(1) diff --git a/datamint/api/endpoints/annotationsets_api.py b/datamint/api/endpoints/annotationsets_api.py index e12f78e9..4aef790e 100644 --- a/datamint/api/endpoints/annotationsets_api.py +++ b/datamint/api/endpoints/annotationsets_api.py @@ -1,8 +1,11 @@ -from ..entity_base_api import CreatableEntityApi, UpdatableEntityApi +import logging from typing import TYPE_CHECKING, Any, Literal, overload -from datamint.entities.annotation_worklist import AnnotationWorklist + from typing_extensions import override -import logging + +from datamint.entities.annotation_worklist import AnnotationWorklist + +from ..entity_base_api import CreatableEntityApi, UpdatableEntityApi if TYPE_CHECKING: from datamint.entities import Project @@ -367,11 +370,11 @@ def upload_annotations( for i, img in enumerate(images): if hasattr(img, 'read'): filename = getattr(img, 'name', f'image_{i}') - files[f'images'] = (filename, img) + files['images'] = (filename, img) else: f = open(img, 'rb') opened_files.append(f) - files[f'images'] = (img, f) + files['images'] = (img, f) return self._make_entity_request('POST', worklist_id, f'resources/{resource_id_str}/segmentations', files=files, data={'payload': payload}).json() diff --git a/datamint/api/endpoints/channels_api.py b/datamint/api/endpoints/channels_api.py index 44cdf5d7..26d66b3f 100644 --- a/datamint/api/endpoints/channels_api.py +++ b/datamint/api/endpoints/channels_api.py @@ -7,10 +7,13 @@ """ import logging + import httpx -from ..entity_base_api import EntityBaseApi + from datamint.entities.channel import Channel +from ..entity_base_api import EntityBaseApi + logger = logging.getLogger(__name__) diff --git a/datamint/api/endpoints/datasetsinfo_api.py b/datamint/api/endpoints/datasetsinfo_api.py index eb501d40..29709305 100644 --- a/datamint/api/endpoints/datasetsinfo_api.py +++ b/datamint/api/endpoints/datasetsinfo_api.py @@ -1,11 +1,13 @@ -from typing import Literal, TYPE_CHECKING from pathlib import Path +from typing import TYPE_CHECKING, Literal -from ..entity_base_api import ApiConfig, DeletableEntityApi -from datamint.entities.datasetinfo import DatasetInfo import httpx from tqdm.auto import tqdm +from datamint.entities.datasetinfo import DatasetInfo + +from ..entity_base_api import ApiConfig, DeletableEntityApi + if TYPE_CHECKING: from datamint.entities import Project diff --git a/datamint/api/endpoints/deploy_model_api.py b/datamint/api/endpoints/deploy_model_api.py index ab7f32db..e7bbce38 100644 --- a/datamint/api/endpoints/deploy_model_api.py +++ b/datamint/api/endpoints/deploy_model_api.py @@ -1,15 +1,16 @@ """API handler for model deployment endpoints.""" -from typing import Any, Literal -from collections.abc import Callable, Generator import json import logging import time +from collections.abc import Callable, Generator +from typing import Any import httpx -from datamint.exceptions import ResourceNotFoundError, JobTimeoutError -from ..entity_base_api import EntityBaseApi, ApiConfig from datamint.entities.deployjob import DeployJob +from datamint.exceptions import JobTimeoutError, ResourceNotFoundError + +from ..entity_base_api import ApiConfig, EntityBaseApi logger = logging.getLogger(__name__) diff --git a/datamint/api/endpoints/inference_api.py b/datamint/api/endpoints/inference_api.py index e5b343ca..c990271d 100644 --- a/datamint/api/endpoints/inference_api.py +++ b/datamint/api/endpoints/inference_api.py @@ -1,22 +1,23 @@ """API handler for model inference endpoints (MLflow DataMint server).""" -from typing import Any, Literal, TYPE_CHECKING -from collections.abc import Callable, Generator import json import logging import time +from collections.abc import Callable, Generator +from typing import TYPE_CHECKING, Any, Literal import httpx -from ..entity_base_api import EntityBaseApi, ApiConfig from datamint.entities.inferencejob import InferenceJob from datamint.exceptions import ( - JobTimeoutError, ItemNotFoundError, - ValidationError, + JobTimeoutError, ModelNotDeployedError, + ValidationError, ) from datamint.mlflow.flavors.model_parser import parse_model_reference +from ..entity_base_api import ApiConfig, EntityBaseApi + if TYPE_CHECKING: from .projects_api import ProjectsApi diff --git a/datamint/api/endpoints/models_api.py b/datamint/api/endpoints/models_api.py index 92c08d9b..e93d4ef9 100644 --- a/datamint/api/endpoints/models_api.py +++ b/datamint/api/endpoints/models_api.py @@ -7,6 +7,7 @@ import mlflow.tracking from datamint.exceptions import ItemNotFoundError + from ..entity_base_api import ApiConfig, BaseApi from .deploy_model_api import DeployModelApi from .model_types import Model, ModelVersion @@ -32,7 +33,6 @@ def __init__(self, @property def _mlflow_client(self) -> mlflow.tracking.MlflowClient: - import datamint.mlflow return mlflow.tracking.MlflowClient() def get_list(self, @@ -178,7 +178,11 @@ def clone_model(self, from datamint.mlflow.flavors import datamint_flavor from datamint.mlflow.flavors.datamint_flavor import FLAVOR_NAME - from datamint.mlflow.tracking.fluent import _reset_active_project, get_active_project_id, set_project + from datamint.mlflow.tracking.fluent import ( + _reset_active_project, + get_active_project_id, + set_project, + ) source_version = self._resolve_version(model, version, alias) diff --git a/datamint/api/endpoints/projects_api.py b/datamint/api/endpoints/projects_api.py index 30819193..4be537c2 100644 --- a/datamint/api/endpoints/projects_api.py +++ b/datamint/api/endpoints/projects_api.py @@ -1,17 +1,25 @@ -from typing import Literal, TYPE_CHECKING, overload from collections.abc import Sequence from pathlib import Path +from typing import TYPE_CHECKING, Literal, overload + +import httpx -from ..entity_base_api import ApiConfig, CRUDEntityApi from datamint.entities.project import Project from datamint.entities.project_resource_split import ProjectResourceSplit -import httpx from datamint.entities.resource import Resource -from datamint.exceptions import DefaultProjectNotSetError, EntityAlreadyExistsError, ItemNotFoundError +from datamint.exceptions import ( + DefaultProjectNotSetError, + EntityAlreadyExistsError, + ItemNotFoundError, +) + +from ..entity_base_api import ApiConfig, CRUDEntityApi + if TYPE_CHECKING: - from . import AnnotationWorklistApi, ResourcesApi from datamint.entities.annotation_worklist import AnnotationWorklist + from . import AnnotationWorklistApi, ResourcesApi + class ProjectsApi(CRUDEntityApi[Project]): """API handler for project-related endpoints.""" @@ -548,11 +556,10 @@ def download_annotations(self, params=params) as response: total_size = int(response.headers.get('content-length', 0)) or None with tqdm(total=total_size, unit='B', unit_scale=True, - disable=not progress_bar) as pbar: - with open(output_path, 'wb') as f: - for chunk in response.iter_bytes(8192): - pbar.update(len(chunk)) - f.write(chunk) + disable=not progress_bar) as pbar, open(output_path, 'wb') as f: + for chunk in response.iter_bytes(8192): + pbar.update(len(chunk)) + f.write(chunk) # ------------------------------------------------------------------ # Statistics diff --git a/datamint/api/endpoints/resources_api.py b/datamint/api/endpoints/resources_api.py index 6f55111b..b791cdc7 100644 --- a/datamint/api/endpoints/resources_api.py +++ b/datamint/api/endpoints/resources_api.py @@ -1,33 +1,35 @@ -from typing import TypeAlias, Literal, IO, overload -from collections.abc import Callable, Sequence -from ..base_api import ApiConfig, BaseApi -from ..entity_base_api import CreatableEntityApi, DeletableEntityApi -from datamint.entities import Project, Resource -from datamint.entities.annotations.annotation import Annotation -from datamint.exceptions import ItemNotFoundError, ServerError, ValidationError -from datamint.entities.annotations import AnnotationType -from datamint.utils.collection_utils import ChainedSequence -import httpx -from datetime import date +import asyncio +import io import json import logging +import os +from collections import defaultdict +from collections.abc import Callable, Sequence +from datetime import date +from pathlib import Path +from typing import IO, Literal, TypeAlias, overload + +import aiohttp +import httpx import pydicom -from pydicom import config as pydicom_config -from medimgkit.dicom_utils import anonymize_dicom, to_bytesio, is_dicom, is_dicom_report from medimgkit import dicom_utils, standardize_mimetype -from medimgkit.io_utils import is_io_object, peek -from medimgkit.format_detection import guess_typez, guess_extension, DEFAULT_MIME_TYPE +from medimgkit.dicom_utils import anonymize_dicom, is_dicom, is_dicom_report, to_bytesio +from medimgkit.format_detection import DEFAULT_MIME_TYPE, guess_extension, guess_typez +from medimgkit.io_utils import is_io_object from medimgkit.nifti_utils import DEFAULT_NIFTI_MIME, NIFTI_MIMES -import os -from tqdm.auto import tqdm -import asyncio -import aiohttp -from pathlib import Path from PIL import Image -import io +from pydicom import config as pydicom_config +from tqdm.auto import tqdm + +from datamint.entities import Project, Resource +from datamint.entities.annotations import AnnotationType +from datamint.entities.annotations.annotation import Annotation +from datamint.exceptions import ItemNotFoundError, ServerError, ValidationError from datamint.types import ImagingData -from collections import defaultdict +from datamint.utils.collection_utils import ChainedSequence +from ..base_api import ApiConfig, BaseApi +from ..entity_base_api import CreatableEntityApi, DeletableEntityApi _LOGGER = logging.getLogger(__name__) _USER_LOGGER = logging.getLogger('user_logger') @@ -64,7 +66,7 @@ def __getattr__(self, attr): def _open_io(file_path: str | Path | IO, mode: str = 'rb') -> IO: - if isinstance(file_path, str) or isinstance(file_path, Path): + if isinstance(file_path, (str, Path)): return open(file_path, 'rb') return file_path @@ -277,7 +279,7 @@ async def _upload_single_resource_async(self, mimetype: str | None = None, anonymize: bool = False, anonymize_retain_codes: Sequence[tuple] = [], - tags: list[str] = [], + tags: list[str] | None = None, mung_filename: Sequence[int] | Literal['all'] | None = None, channel: str | None = None, session=None, @@ -285,6 +287,8 @@ async def _upload_single_resource_async(self, publish: bool = False, metadata_file: str | dict | None = None, ) -> str: + if tags is None: + tags = [] if is_io_object(file_path): source_filepath = os.path.abspath(os.path.expanduser(file_path.name)) filename = os.path.basename(source_filepath) @@ -328,7 +332,7 @@ async def _upload_single_resource_async(self, mimetype = standardize_mimetype(mimetype) - if is_a_dicom_file == True or is_dicom(file_path): + if is_a_dicom_file or is_dicom(file_path): if tags is None: tags = [] else: @@ -414,7 +418,7 @@ async def _upload_single_resource_async(self, form.add_field('bypass_inbox', 'true' if publish else 'false') if tags is not None and len(tags) > 0: # comma separated list of tags - form.add_field('tags', ','.join([l.strip() for l in tags])) + form.add_field('tags', ','.join([tag.strip() for tag in tags])) # Add JSON metadata if provided if metadata_content is not None: @@ -545,7 +549,7 @@ async def __upload_single_resource(all_files_path, index: int, tasks = [__upload_single_resource(files_path, i, segfiles, metadata_file) for i, segfiles, metadata_file in zip(range(len(files_path)), segmentation_files, metadata_files)] except ValueError: - msg = f"Error preparing upload tasks. Try `assemble_dicom=False`." + msg = "Error preparing upload tasks. Try `assemble_dicom=False`." _LOGGER.error(msg) _USER_LOGGER.error(msg) raise @@ -565,8 +569,7 @@ def _validate_upload_params( if on_error not in ['raise', 'skip']: raise ValueError("on_error must be either 'raise' or 'skip'") if ( - isinstance(files_path, IO) - or isinstance(files_path, pydicom.Dataset) + isinstance(files_path, (IO, pydicom.Dataset)) or (isinstance(files_path, str) and not os.path.isdir(files_path)) ): raise ValueError( @@ -680,7 +683,7 @@ def _post_upload_add_to_project( except Exception as e: _LOGGER.error(f"Error adding resources to project: {e}") if on_error == 'raise': - raise e + raise def upload_resources(self, files_path: Sequence[str | IO | pydicom.Dataset], @@ -973,7 +976,7 @@ async def _async_download_file(self, except ItemNotFoundError as e: e.set_params('resource', {'resource_id': resource_id}) - raise e + raise def download_multiple_resources(self, resources: Sequence[str] | Sequence[Resource], @@ -1128,7 +1131,7 @@ def download_resource_file(self, except ValueError as e: _LOGGER.warning(f"Could not convert file to a known format: {e}") resource_file = response.content - except NotImplementedError as e: + except NotImplementedError: _LOGGER.warning(f"Conversion not implemented yet for {mimetype} and save_path=None." + " Returning a bytes array. If you want the conversion for this mimetype, provide a save_path.") resource_file = response.content @@ -1136,7 +1139,7 @@ def download_resource_file(self, resource_file = response.content except ItemNotFoundError as e: e.set_params('resource', {'resource_id': self._entid(resource)}) - raise e + raise if save_path is not None: if add_extension and mimetype is not None: @@ -1239,7 +1242,7 @@ def download_resource_frame(self, status_code=response.status_code) except ItemNotFoundError as e: e.set_params('resource', {'resource_id': self._entid(resource)}) - raise e + raise def publish_resources(self, resources: str | Resource | Sequence[str | Resource]) -> None: @@ -1356,7 +1359,7 @@ def bulk_delete(self, entities: Sequence[str | Resource]) -> None: if len(resources_ids) == 0: return batch_size = 200 - for i in range(0, ceil(len(resources_ids)/batch_size)): + for i in range(ceil(len(resources_ids)/batch_size)): batch_ids = resources_ids[i*batch_size:(i+1)*batch_size] self._make_request('DELETE', f'{self.endpoint_base}', diff --git a/datamint/api/endpoints/users_api.py b/datamint/api/endpoints/users_api.py index ad081fa5..9333314c 100644 --- a/datamint/api/endpoints/users_api.py +++ b/datamint/api/endpoints/users_api.py @@ -1,9 +1,11 @@ -from typing import Literal, TYPE_CHECKING, cast, overload +from typing import TYPE_CHECKING, Literal, cast, overload -from ..entity_base_api import CreatableEntityApi, ApiConfig -from datamint.entities import User import httpx +from datamint.entities import User + +from ..entity_base_api import ApiConfig, CreatableEntityApi + if TYPE_CHECKING: from datamint.entities import Project @@ -70,13 +72,13 @@ def create(self, The created user entity or identifier, depending on ``return_entity``. The identifier is the user's email address. """ - data = dict( - email=email, - password=password, - firstname=firstname, - lastname=lastname, - roles=roles - ) + data = { + 'email': email, + 'password': password, + 'firstname': firstname, + 'lastname': lastname, + 'roles': roles + } return cast(str | User | None, self._create(data, return_entity=return_entity, exists_ok=exists_ok)) def get_by_id(self, entity_id: str) -> User: diff --git a/datamint/api/entity_base_api.py b/datamint/api/entity_base_api.py index 10461802..c4e4722c 100644 --- a/datamint/api/entity_base_api.py +++ b/datamint/api/entity_base_api.py @@ -1,14 +1,17 @@ -from typing import Any, Literal, TypeVar, Generic, overload -from collections.abc import Sequence, AsyncGenerator +import asyncio +import contextlib import logging +import re +from collections.abc import AsyncGenerator, Sequence +from typing import Any, Generic, Literal, TypeVar, overload + +import aiohttp import httpx + from datamint.entities.base_entity import BaseEntity from datamint.exceptions import ItemNotFoundError, ServerError -import aiohttp -import asyncio + from .base_api import ApiConfig, BaseApi -import contextlib -import re _UUID_PATTERN = re.compile( r'^[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}$', @@ -368,7 +371,7 @@ def _create(self, entity_data: dict[str, Any], if return_entity: try: return self._init_entity_obj(**respdata) - except: + except Exception: logger.debug("Failed to init entity obj on create response. Falling back to get_by_id.") return self.get_by_id(respdata.get('id')) return respdata.get('id') @@ -425,4 +428,3 @@ def partial_update(self, entity: str | T, entity_data: dict[str, Any]): class CRUDEntityApi(CreatableEntityApi[T], UpdatableEntityApi[T], DeletableEntityApi[T]): """Full CRUD API handler for entities supporting create, read, update, delete operations.""" - pass diff --git a/datamint/client_cmd_tools/datamint_config.py b/datamint/client_cmd_tools/datamint_config.py index dad22227..3b0b9097 100644 --- a/datamint/client_cmd_tools/datamint_config.py +++ b/datamint/client_cmd_tools/datamint_config.py @@ -11,7 +11,10 @@ from typing_extensions import NotRequired from datamint import configs -from datamint.utils.logging_utils import ConsoleWrapperHandler, load_cmdline_logging_config +from datamint.utils.logging_utils import ( + ConsoleWrapperHandler, + load_cmdline_logging_config, +) _LOGGER = logging.getLogger(__name__) _USER_LOGGER = logging.getLogger('user_logger') @@ -49,7 +52,7 @@ def configure_default_url(): return # Basic URL validation - if not (url.startswith('http://') or url.startswith('https://')): + if not (url.startswith(('http://', 'https://'))): console.print("[warning]⚠️ URL should start with http:// or https://[/warning]") return @@ -118,7 +121,7 @@ def test_connection(): from datamint import Api console.print("[accent]🔄 Testing connection...[/accent]") Api(check_connection=True) - console.print(f"[success]✅ Connection successful![/success]") + console.print("[success]✅ Connection successful![/success]") except ImportError: console.print("[error]❌ Full API not available. Install with: pip install datamint[/error]") except Exception as e: @@ -753,9 +756,9 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa When ``subparsers`` is given, the parser is registered as a ``config`` subparser (used by ``datamint``'s combined completion tree) instead of a standalone parser. """ - kwargs = dict( - description='🔧 Datamint API Configuration Tool', - epilog=""" + kwargs = { + 'description': '🔧 Datamint API Configuration Tool', + 'epilog': """ Examples: datamint config # Interactive mode datamint config --api-key YOUR_KEY # Set API key @@ -773,8 +776,8 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa More Documentation: https://sonanceai.github.io/datamint-python-api/command_line_tools.html """, - formatter_class=argparse.RawDescriptionHelpFormatter - ) + 'formatter_class': argparse.RawDescriptionHelpFormatter + } if subparsers is not None: parser = subparsers.add_parser('config', **kwargs) else: diff --git a/datamint/client_cmd_tools/datamint_example.py b/datamint/client_cmd_tools/datamint_example.py index 5dfafe55..75c37c3f 100644 --- a/datamint/client_cmd_tools/datamint_example.py +++ b/datamint/client_cmd_tools/datamint_example.py @@ -24,9 +24,9 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa When ``subparsers`` is given, the parser is registered as an ``example`` subparser (used by ``datamint``'s combined completion tree) instead of a standalone parser. """ - kwargs = dict( - description='Populate a Datamint project with an example dataset.', - epilog=""" + kwargs = { + 'description': 'Populate a Datamint project with an example dataset.', + 'epilog': """ Examples: datamint example bccd # Blood cell detection (BCCD) datamint example busi --project MyBusiProject # Breast ultrasound segmentation (BUSI) @@ -35,8 +35,8 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa More Documentation: https://sonanceai.github.io/datamint-python-api/command_line_tools.html """, - formatter_class=argparse.RawDescriptionHelpFormatter, - ) + 'formatter_class': argparse.RawDescriptionHelpFormatter, + } if subparsers is not None: parser = subparsers.add_parser('example', **kwargs) else: diff --git a/datamint/client_cmd_tools/datamint_inference.py b/datamint/client_cmd_tools/datamint_inference.py index 6b74fcfe..f2b2bb72 100644 --- a/datamint/client_cmd_tools/datamint_inference.py +++ b/datamint/client_cmd_tools/datamint_inference.py @@ -14,9 +14,15 @@ from rich.console import Console from rich.table import Table -from datamint.client_cmd_tools.datamint_upload import _is_valid_path_argparse, handle_api_key +from datamint.client_cmd_tools.datamint_upload import ( + _is_valid_path_argparse, + handle_api_key, +) from datamint.exceptions import DatamintException, ItemNotFoundError -from datamint.utils.logging_utils import ConsoleWrapperHandler, load_cmdline_logging_config +from datamint.utils.logging_utils import ( + ConsoleWrapperHandler, + load_cmdline_logging_config, +) _LOGGER = logging.getLogger(__name__) _USER_LOGGER = logging.getLogger('user_logger') @@ -88,9 +94,9 @@ def _print_predictions(console: Console, predictions: list) -> None: def _save_overlay(console: Console, resource: Any, predictions: list, output_path: str) -> None: import matplotlib matplotlib.use('Agg') - import matplotlib.patches as patches import matplotlib.pyplot as plt import numpy as np + from matplotlib import patches img_np = np.array(resource.fetch_file_data(auto_convert=True)) if img_np.ndim == 3 and img_np.shape[0] in (1, 3) and img_np.shape[0] != img_np.shape[-1]: @@ -163,9 +169,9 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa When ``subparsers`` is given, the parser is registered as an ``inference`` subparser (used by ``datamint``'s combined completion tree) instead of a standalone parser. """ - kwargs = dict( - description='Run local inference with a registered Datamint model against a local file.', - epilog=""" + kwargs = { + 'description': 'Run local inference with a registered Datamint model against a local file.', + 'epilog': """ Examples: datamint inference file.png --model-name MyModel # Predict using the model registered as 'MyModel' @@ -178,8 +184,8 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa More Documentation: https://sonanceai.github.io/datamint-python-api/command_line_tools.html """, - formatter_class=argparse.RawDescriptionHelpFormatter, - ) + 'formatter_class': argparse.RawDescriptionHelpFormatter, + } if subparsers is not None: parser = subparsers.add_parser('inference', **kwargs) else: @@ -212,7 +218,7 @@ def _parse_args() -> argparse.Namespace: def main() -> None: global CONSOLE load_cmdline_logging_config() - CONSOLE = [h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)][0].console + CONSOLE = next(h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)).console args = _parse_args() diff --git a/datamint/client_cmd_tools/datamint_init.py b/datamint/client_cmd_tools/datamint_init.py index 8fcd1ce4..ef3cf8a9 100644 --- a/datamint/client_cmd_tools/datamint_init.py +++ b/datamint/client_cmd_tools/datamint_init.py @@ -2,10 +2,9 @@ from pathlib import Path from rich.console import Console -from rich.prompt import Prompt, Confirm +from rich.prompt import Confirm, Prompt from rich.rule import Rule - console = Console() # --------------------------------------------------------------------------- diff --git a/datamint/client_cmd_tools/datamint_train.py b/datamint/client_cmd_tools/datamint_train.py index 7d3f82f9..78892be3 100644 --- a/datamint/client_cmd_tools/datamint_train.py +++ b/datamint/client_cmd_tools/datamint_train.py @@ -14,7 +14,7 @@ import logging import os import sys -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any from rich.console import Console from rich.prompt import Confirm, Prompt @@ -23,7 +23,10 @@ from datamint import Api, configs from datamint.client_cmd_tools.datamint_upload import handle_api_key from datamint.exceptions import DatamintException -from datamint.utils.logging_utils import ConsoleWrapperHandler, load_cmdline_logging_config +from datamint.utils.logging_utils import ( + ConsoleWrapperHandler, + load_cmdline_logging_config, +) if TYPE_CHECKING: from datamint.dataset.base import DatamintBaseDataset @@ -78,13 +81,17 @@ class DatamintTrainCliError(Exception): """Expected, user-facing CLI error (bad input, ambiguous project, etc.).""" -def _detect(console: Console, project_name: str) -> tuple['DatamintBaseDataset', str, list[str]]: +def _detect(console: Console, project_name: str) -> tuple[DatamintBaseDataset, str, list[str]]: """Auto-detect data format and candidate task(s) for a project. Returns (dataset, format, tasks_present) where format is '2d' or '3d' and tasks_present is the list of task names whose annotations were found (usually one). """ - from datamint.dataset import ImageDataset, VideoDataset, VolumeDataset, build_dataset + from datamint.dataset import ( + ImageDataset, + VideoDataset, + build_dataset, + ) with console.status("[accent]Detecting task and data format...[/accent]"): try: @@ -190,8 +197,8 @@ def _build_trainer_kwargs(args: argparse.Namespace, alias: str) -> dict[str, Any def _print_plan(console: Console, - project: 'Project', - dataset: 'DatamintBaseDataset', + project: Project, + dataset: DatamintBaseDataset, fmt: str, task: str, model_alias: str, @@ -238,7 +245,7 @@ def _print_results(console: Console, trainer, results: dict[str, Any]) -> None: console.print(f"MLflow experiment: [key]{trainer.experiment_name}[/key]") -def _resolve_project(api: Api, name: str) -> 'Project': +def _resolve_project(api: Api, name: str) -> Project: project = api.projects.get_by_name(name) if project is not None: return project @@ -356,9 +363,9 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa When ``subparsers`` is given, the parser is registered as a ``train`` subparser (used by ``datamint``'s combined completion tree) instead of a standalone parser. """ - kwargs = dict( - description='Train a model on a Datamint project using a built-in one-line trainer.', - epilog=""" + kwargs = { + 'description': 'Train a model on a Datamint project using a built-in one-line trainer.', + 'epilog': """ Examples: datamint train --project MyProject --model yolox --max-epochs 20 # Train a specific model @@ -369,8 +376,8 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa More Documentation: https://sonanceai.github.io/datamint-python-api/command_line_tools.html#training-a-model """, - formatter_class=argparse.RawDescriptionHelpFormatter, - ) + 'formatter_class': argparse.RawDescriptionHelpFormatter, + } if subparsers is not None: parser = subparsers.add_parser('train', **kwargs) else: @@ -413,7 +420,7 @@ def _parse_args() -> argparse.Namespace: def main() -> None: global CONSOLE load_cmdline_logging_config() - CONSOLE = [h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)][0].console + CONSOLE = next(h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)).console args = _parse_args() diff --git a/datamint/client_cmd_tools/datamint_upload.py b/datamint/client_cmd_tools/datamint_upload.py index 04e2f24d..89feec7e 100644 --- a/datamint/client_cmd_tools/datamint_upload.py +++ b/datamint/client_cmd_tools/datamint_upload.py @@ -1,24 +1,27 @@ -from datamint.exceptions import DatamintException, ItemNotFoundError import argparse -from datamint import Api -import os -from humanize import naturalsize +import fnmatch import logging -from pathlib import Path +import os import sys -from medimgkit.dicom_utils import is_dicom, detect_dicomdir, parse_dicomdir_files -import fnmatch -from typing import Any -from collections.abc import Generator from collections import defaultdict -from datamint import __version__ as datamint_version -from datamint import configs -from datamint.utils.logging_utils import load_cmdline_logging_config, ConsoleWrapperHandler -from rich.console import Console -import yaml -from collections.abc import Iterable +from collections.abc import Generator, Iterable +from pathlib import Path +from typing import Any + import pandas as pd import pydicom.errors +import yaml +from humanize import naturalsize +from medimgkit.dicom_utils import detect_dicomdir, is_dicom, parse_dicomdir_files +from rich.console import Console + +from datamint import Api, configs +from datamint import __version__ as datamint_version +from datamint.exceptions import DatamintException, ItemNotFoundError +from datamint.utils.logging_utils import ( + ConsoleWrapperHandler, + load_cmdline_logging_config, +) # Create two loggings: one for the user and one for the developer _LOGGER = logging.getLogger(__name__) @@ -132,7 +135,7 @@ def _is_valid_path_argparse(x): argparse type that checks if the path exists """ if not os.path.exists(x): - raise argparse.ArgumentTypeError("{0} does not exist".format(x)) + raise argparse.ArgumentTypeError(f"{x} does not exist") return x @@ -203,8 +206,7 @@ def walk_to_depth(path: str | Path, _USER_LOGGER.info(f"Found DICOMDIR file at {path}. Using it as authoritative source for file listing.") dicom_files = parse_dicomdir_files(dicomdir_path) # Yield all DICOM files from DICOMDIR and return early - for dicom_file in dicom_files: - yield dicom_file + yield from dicom_files return except Exception as e: _USER_LOGGER.warning(f"Failed to parse DICOMDIR at {path}: {e}. Falling back to directory scan.") @@ -524,8 +526,8 @@ def _build_parser(subparsers: argparse._SubParsersAction | None = None) -> argpa When ``subparsers`` is given, the parser is registered as an ``upload`` subparser (used by ``datamint``'s combined completion tree) instead of a standalone parser. """ - kwargs = dict( - description='DatamintAPI command line tool for uploading DICOM files and other resources') + kwargs = { + 'description': 'DatamintAPI command line tool for uploading DICOM files and other resources'} if subparsers is not None: parser = subparsers.add_parser('upload', **kwargs) else: @@ -693,7 +695,7 @@ def _parse_args() -> tuple[Any, list[str], list[dict] | None, list[str] | None]: except Exception as e: if args.verbose: _LOGGER.exception(e) - raise e + raise def print_input_summary(files_path: list[str], @@ -724,7 +726,7 @@ def print_input_summary(files_path: list[str], _USER_LOGGER.info("\t(...)") _USER_LOGGER.info(f"\t{distinguishing_paths[files_path[-1]]}") _USER_LOGGER.info(f"Total size of the upload: {naturalsize(total_size)}") - _USER_LOGGER.info(f"Number of files per extension:") + _USER_LOGGER.info("Number of files per extension:") for ext, count in ext_counts: if ext == '': ext = 'no extension' @@ -736,7 +738,7 @@ def print_input_summary(files_path: list[str], if segfiles is not None: num_segfiles = sum([1 if seg is not None else 0 for seg in segfiles]) - msg = f"Number of images with an associated segmentation: " +\ + msg = "Number of images with an associated segmentation: " +\ f"{num_segfiles} ({num_segfiles / total_files:.0%})" if num_segfiles == 0: _USER_LOGGER.warning(msg) @@ -745,7 +747,7 @@ def print_input_summary(files_path: list[str], # count number of segmentations files with names if args.segmentation_names is not None and num_segfiles > 0: segnames_count = sum([1 if 'names' in seg else 0 for seg in segfiles if seg is not None]) - msg = f"Number of segmentations with associated name: " + \ + msg = "Number of segmentations with associated name: " + \ f"{segnames_count} ({segnames_count / num_segfiles:.0%})" if segnames_count == 0: _USER_LOGGER.warning(msg) @@ -767,25 +769,25 @@ def print_results_summary(files_path: list[str], # Get distinguishing paths for better error reporting distinguishing_paths = _get_minimal_distinguishing_paths(files_path) - _USER_LOGGER.info(f"\nUpload summary:") + _USER_LOGGER.info("\nUpload summary:") _USER_LOGGER.info(f"\tTotal files: {len(files_path)}") _USER_LOGGER.info(f"\tSuccessful uploads: {len(files_path) - len(failure_files)}") if len(failure_files) > 0: _USER_LOGGER.warning(f"\tFailed uploads: {len(failure_files)}") _USER_LOGGER.warning(f"\tFailed files: {[distinguishing_paths[f] for f in failure_files]}") - _USER_LOGGER.warning(f"\nFailures:") + _USER_LOGGER.warning("\nFailures:") for f, r in zip(files_path, results): if isinstance(r, Exception): _USER_LOGGER.warning(f"\t{distinguishing_paths[f]}: {r}") else: - CONSOLE.print(f'✅ All uploads successful!', style='success') + CONSOLE.print('✅ All uploads successful!', style='success') return len(failure_files) def main(): global CONSOLE load_cmdline_logging_config() - CONSOLE = [h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)][0].console + CONSOLE = next(h for h in _USER_LOGGER.handlers if isinstance(h, ConsoleWrapperHandler)).console try: args, files_path, segfiles, metadata_files = _parse_args() @@ -824,7 +826,7 @@ def main(): files_path=files_path, tags=args.tag, on_error='skip', - anonymize=args.retain_pii == False and has_a_dicom_file, + anonymize=not args.retain_pii and has_a_dicom_file, anonymize_retain_codes=args.retain_attribute, mung_filename=args.mungfilename, publish=args.publish, diff --git a/datamint/configs.py b/datamint/configs.py index 3cba9a0d..db41b2c8 100644 --- a/datamint/configs.py +++ b/datamint/configs.py @@ -1,9 +1,10 @@ -import yaml -import os import logging -from platformdirs import PlatformDirs +import os from typing import Any +import yaml +from platformdirs import PlatformDirs + APIURL_KEY = 'default_api_url' APIKEY_KEY = 'api_key' @@ -59,11 +60,10 @@ def set_values(values: dict[str, Any]): def get_value(key: str, include_envvars: bool = True): - if include_envvars: - if key in ENV_VARS: - env_var = os.getenv(ENV_VARS[key]) - if env_var is not None: - return env_var + if include_envvars and key in ENV_VARS: + env_var = os.getenv(ENV_VARS[key]) + if env_var is not None: + return env_var config = read_config() return config.get(key) diff --git a/datamint/dataset/__init__.py b/datamint/dataset/__init__.py index 1cfe80ea..345a73ee 100644 --- a/datamint/dataset/__init__.py +++ b/datamint/dataset/__init__.py @@ -11,29 +11,28 @@ # New modular architecture from .base import DatamintBaseDataset, DatamintDatasetException +from .factory import build_dataset +from .image_dataset import ImageDataset, detection_collate_fn from .multiframe_dataset import MultiFrameDataset -from .image_dataset import ImageDataset -from .volume_dataset import VolumeDataset -from .video_dataset import VideoDataset from .sliced_dataset import SlicedVolumeDataset from .sliced_video_dataset import SlicedVideoDataset -from .image_dataset import detection_collate_fn -from .factory import build_dataset from .split_result import SplitResult +from .video_dataset import VideoDataset +from .volume_dataset import VolumeDataset __all__ = [ # Core 'DatamintBaseDataset', 'DatamintDatasetException', - 'MultiFrameDataset', # Specialized datasets 'ImageDataset', - 'VolumeDataset', - 'VideoDataset', - 'SlicedVolumeDataset', + 'MultiFrameDataset', 'SlicedVideoDataset', - 'detection_collate_fn', + 'SlicedVolumeDataset', + 'SplitResult', + 'VideoDataset', + 'VolumeDataset', # Factory 'build_dataset', - 'SplitResult', + 'detection_collate_fn', ] \ No newline at end of file diff --git a/datamint/dataset/annotation_processor.py b/datamint/dataset/annotation_processor.py index bc93a7f7..a6c883c3 100644 --- a/datamint/dataset/annotation_processor.py +++ b/datamint/dataset/annotation_processor.py @@ -10,17 +10,16 @@ for any dataset type, while specialized logic is in subclasses. """ import logging -from typing import Literal, TYPE_CHECKING -from collections.abc import Iterable, Sequence from collections import defaultdict +from collections.abc import Iterable, Sequence +from typing import TYPE_CHECKING, Literal import numpy as np import torch +from medimgkit.readers import read_array_normalized from torch import Tensor from typing_extensions import overload -from medimgkit.readers import read_array_normalized - if TYPE_CHECKING: from datamint.entities.annotations.annotation import Annotation @@ -266,7 +265,7 @@ def load_segmentations( """ seg_annotations = [ann for ann in annotations if ann.annotation_type == 'segmentation'] - uniq_authors = set(ann.created_by or ann.created_by_model or "unknown" for ann in seg_annotations) + uniq_authors = {ann.created_by or ann.created_by_model or "unknown" for ann in seg_annotations} segmentations: dict[str, list[np.ndarray]] = {a: [] for a in uniq_authors} # tensors of shape (D, H, W) seg_labels: dict[str, list[int]] = {a: [] for a in uniq_authors} # list of size=#num_instances @@ -337,21 +336,21 @@ def load_segmentation_data(self, ann: 'Annotation', def _merge_union(self, segmentations: dict[str, Tensor]) -> Tensor: """Union merge: pixel is labeled if ANY annotator labeled it.""" - new_segmentations = torch.zeros_like(list(segmentations.values())[0]) + new_segmentations = torch.zeros_like(next(iter(segmentations.values()))) for seg in segmentations.values(): new_segmentations += seg return new_segmentations.bool().to(torch.uint8) def _merge_intersection(self, segmentations: dict[str, Tensor]) -> Tensor: """Intersection merge: pixel is labeled if ALL annotators labeled it.""" - new_segmentations = torch.ones_like(list(segmentations.values())[0]) + new_segmentations = torch.ones_like(next(iter(segmentations.values()))) for seg in segmentations.values(): new_segmentations *= seg return new_segmentations.bool().to(torch.uint8) def _merge_mode(self, segmentations: dict[str, Tensor]) -> Tensor: """Mode merge: pixel is labeled if majority of annotators labeled it.""" - new_segmentations = torch.zeros_like(list(segmentations.values())[0]) + new_segmentations = torch.zeros_like(next(iter(segmentations.values()))) for seg in segmentations.values(): new_segmentations += seg new_segmentations = new_segmentations >= len(segmentations) / 2 diff --git a/datamint/dataset/base.py b/datamint/dataset/base.py index 31bf7985..ac2ef95f 100644 --- a/datamint/dataset/base.py +++ b/datamint/dataset/base.py @@ -6,26 +6,32 @@ """ import logging from abc import ABC, abstractmethod +from collections.abc import Callable, Iterator, Sequence from datetime import datetime, timezone -from typing import Any, TYPE_CHECKING, Literal, cast -from collections.abc import Sequence, Callable, Iterator +from typing import TYPE_CHECKING, Any, Literal, cast +import numpy as np import torch from torch import Tensor -from torch.utils.data import DataLoader, ConcatDataset -import numpy as np +from torch.utils.data import ConcatDataset, DataLoader + +from datamint._repr_utils import render_html_card, render_text_block from datamint.entities.annotation_worklist import AnnotationWorklist -from datamint.exceptions import DatamintException, ItemNotFoundError -from datamint._repr_utils import render_text_block, render_html_card -from .annotation_processor import AnnotationProcessor, MergeStrategy -from datamint.entities.annotations.annotation_spec import AnnotationSpec, CategoryAnnotationSpec from datamint.entities.annotations import AnnotationType +from datamint.entities.annotations.annotation_spec import ( + AnnotationSpec, + CategoryAnnotationSpec, +) +from datamint.exceptions import DatamintException, ItemNotFoundError +from .annotation_processor import AnnotationProcessor, MergeStrategy if TYPE_CHECKING: - from datamint.entities import Resource, Project, Annotation from albumentations import BaseCompose + + from datamint.entities import Annotation, Project, Resource from datamint.mlflow.data import DatamintMLflowDataset + from .split_result import SplitResult _LOGGER = logging.getLogger(__name__) @@ -33,7 +39,6 @@ class DatamintDatasetException(DatamintException): """Exception raised for dataset errors.""" - pass class DatamintBaseDataset(ABC, torch.utils.data.Dataset): @@ -663,8 +668,8 @@ def _augment_labels_from_annotations(self) -> None: Scans resource annotations for identifiers not present in the project's annotations_specs and adds them to the corresponding label/segmentation mappings. """ - inferred_frame_lsets, inferred_frame_lcodes = self._infer_labels_set(framed=True) - inferred_image_lsets, inferred_image_lcodes = self._infer_labels_set(framed=False) + inferred_frame_lsets, _inferred_frame_lcodes = self._infer_labels_set(framed=True) + inferred_image_lsets, _inferred_image_lcodes = self._infer_labels_set(framed=False) inferred_seglabel_list, _ = self._infer_segmentation_group() # Augment frame labels @@ -722,7 +727,6 @@ def _get_raw_item(self, index: int) -> dict[str, Any]: - 'metainfo': dict - 'annotations': list[Annotation] """ - pass def get_resource(self, index: int) -> 'Resource': """Get the Resource object for a given index.""" @@ -760,12 +764,11 @@ def _should_include_annotation(self, ann: 'Annotation') -> bool: return self._should_include_image_label(ann.identifier) else: # frame-level return self._should_include_frame_label(ann.identifier) - elif ann.is_category(): - if not self.allow_external_annotations: - lsets = self.image_lsets if ann.frame_index is None else self.frame_lsets - valid_identifiers = {ident for ident, _ in lsets.get('multiclass', [])} - if ann.identifier not in valid_identifiers: - return False + elif ann.is_category() and not self.allow_external_annotations: + lsets = self.image_lsets if ann.frame_index is None else self.frame_lsets + valid_identifiers = {ident for ident, _ in lsets.get('multiclass', [])} + if ann.identifier not in valid_identifiers: + return False return True @@ -1123,7 +1126,6 @@ def apply_alb_transform( Returns: Dict with ``'image'`` key plus the same target keys, all transformed. """ - pass def __len__(self) -> int: """Dataset length.""" @@ -1482,11 +1484,11 @@ def _group_resources_indices_by_patient( if pid is None: if none_patient_id_strategy == 'error': - raise ValueError(( + raise ValueError( f"Resource at index {idx} (id={getattr(resource, 'id', '?')!r}) has no patient_id." "Set none_patient_id_strategy='individual' to treat each as its own patient, " "'group' to group all together, or 'skip' to exclude them." - )) + ) elif none_patient_id_strategy == 'skip': continue elif none_patient_id_strategy == 'individual': @@ -1550,7 +1552,7 @@ def _split_locally_by_patient( patient_indices = self._group_resources_indices_by_patient(none_patient_id_strategy) patients_ids = list(patient_indices.keys()) - import random + import random rng = random.Random(seed) rng.shuffle(patients_ids) diff --git a/datamint/dataset/factory.py b/datamint/dataset/factory.py index 8c5d652d..1539f0a9 100644 --- a/datamint/dataset/factory.py +++ b/datamint/dataset/factory.py @@ -4,14 +4,15 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from datamint.entities.resource import Resource from datamint.entities import Project + from datamint.entities.resource import Resource + from .base import DatamintBaseDataset _LOGGER = logging.getLogger(__name__) -def _classify_resource(resource: 'Resource') -> str: +def _classify_resource(resource: Resource) -> str: if resource.is_video(): return 'video' if resource.is_volume(): @@ -27,7 +28,7 @@ def _classify_resource(resource: 'Resource') -> str: def build_dataset(project: str | Project | None = None, - **kwargs: Any) -> 'DatamintBaseDataset': + **kwargs: Any) -> DatamintBaseDataset: """Auto-detect and return the appropriate dataset class for a project. Fetches a small sample of resources from the project to determine the @@ -59,9 +60,10 @@ def build_dataset(project: str | Project | None = None, from datamint import Api from datamint.entities import Project as ProjectCls + from .image_dataset import ImageDataset - from .volume_dataset import VolumeDataset from .video_dataset import VideoDataset + from .volume_dataset import VolumeDataset _KIND_TO_CLS = { 'image': ImageDataset, diff --git a/datamint/dataset/image_dataset.py b/datamint/dataset/image_dataset.py index 75c47d58..4908e874 100644 --- a/datamint/dataset/image_dataset.py +++ b/datamint/dataset/image_dataset.py @@ -6,11 +6,11 @@ """ import logging from typing import Any -from typing_extensions import override + +import numpy as np import torch from torch import Tensor -import numpy as np -import albumentations +from typing_extensions import override from .base import DatamintDatasetException from .volume_dataset import VolumeDataset diff --git a/datamint/dataset/multiframe_dataset.py b/datamint/dataset/multiframe_dataset.py index 2c8a67af..f080debe 100644 --- a/datamint/dataset/multiframe_dataset.py +++ b/datamint/dataset/multiframe_dataset.py @@ -7,13 +7,13 @@ """ import logging from typing import Any -from typing_extensions import override -import torch -import numpy as np import albumentations - +import numpy as np +import torch from medimgkit.readers import read_array_normalized +from typing_extensions import override + from .base import DatamintBaseDataset _LOGGER = logging.getLogger(__name__) diff --git a/datamint/dataset/sliced_dataset.py b/datamint/dataset/sliced_dataset.py index 38eb0fbc..543399df 100644 --- a/datamint/dataset/sliced_dataset.py +++ b/datamint/dataset/sliced_dataset.py @@ -5,25 +5,28 @@ enabling training of 2D models on volumetric medical imaging data. """ from __future__ import annotations + import hashlib -from typing import Any, TYPE_CHECKING, cast -from typing_extensions import override +import logging from collections.abc import Sequence +from typing import TYPE_CHECKING, Any + +import albumentations import numpy as np import torch from torch import Tensor -import albumentations +from typing_extensions import override -from .base import DatamintBaseDataset -from .annotation_processor import AnnotationProcessor +from datamint.entities.cache_manager import CacheManager -import logging +from .annotation_processor import AnnotationProcessor +from .base import DatamintBaseDataset -from datamint.entities.cache_manager import CacheManager if TYPE_CHECKING: + from medimgkit import ViewPlane + from datamint.entities import Annotation, Resource from datamint.entities.sliced_resource import SlicedVolumeResource - from medimgkit import ViewPlane _LOGGER = logging.getLogger(__name__) @@ -143,8 +146,8 @@ def _prepare(self) -> None: def from_dataset( cls, parent_dataset: DatamintBaseDataset, - slice_axis: 'ViewPlane | int' = 'axial', - ) -> 'SlicedVolumeDataset': + slice_axis: ViewPlane | int = 'axial', + ) -> SlicedVolumeDataset: """Create a SlicedVolumeDataset from an existing dataset without additional server calls. Copies all configuration, label mappings, and already-loaded resources @@ -168,7 +171,7 @@ def from_dataset( ) @staticmethod - def _validate_slice_axis(slice_axis: 'ViewPlane | int') -> 'ViewPlane': + def _validate_slice_axis(slice_axis: ViewPlane | int) -> ViewPlane: if isinstance(slice_axis, str): valid_slice_axis = ['axial', 'coronal', 'sagittal'] if slice_axis not in valid_slice_axis: @@ -186,10 +189,10 @@ def _validate_slice_axis(slice_axis: 'ViewPlane | int') -> 'ViewPlane': def _expand_resources( self, - resources: Sequence['Resource'], - resource_annotations: Sequence[Sequence['Annotation']], + resources: Sequence[Resource], + resource_annotations: Sequence[Sequence[Annotation]], volume_cache: CacheManager, - ) -> tuple[list[SlicedVolumeResource], list[Sequence['Annotation']]]: + ) -> tuple[list[SlicedVolumeResource], list[Sequence[Annotation]]]: """Expand volume resources into per-slice proxy resources. Args: @@ -203,7 +206,7 @@ def _expand_resources( from datamint.entities.sliced_resource import SlicedVolumeResource sliced_resources: list[SlicedVolumeResource] = [] - sliced_annotations: list[Sequence['Annotation']] = [] + sliced_annotations: list[Sequence[Annotation]] = [] requires_download = any(not r.is_cached() for r in resources) iterator = enumerate(resources) @@ -242,7 +245,7 @@ def _std_axis_to_seg_axis(slice_axis_idx_std: int) -> int: def _load_sliced_segmentations( self, - annotations: Sequence['Annotation'], + annotations: Sequence[Annotation], resource: SlicedVolumeResource, ) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray], dict[str, list]]: """Load segmentations already sliced for a specific 2D slice, with caching. @@ -270,9 +273,9 @@ def _load_sliced_segmentations( image_seg_anns = [a for a in seg_anns if a.scope == 'image'] frame_seg_anns = [a for a in seg_anns if a.scope == 'frame'] - uniq_authors = set( + uniq_authors = { self.annotation_processor.get_author(a) for a in seg_anns - ) + } segmentations: dict[str, list[np.ndarray]] = {a: [] for a in uniq_authors} seg_labels: dict[str, list[int]] = {a: [] for a in uniq_authors} seg_metainfos: dict[str, list] = {a: [] for a in uniq_authors} @@ -362,7 +365,7 @@ def _fetch_sliced_seg_annotation( def _fetch_sliced_frame_seg_group( self, - fr_anns: list['Annotation'], + fr_anns: list[Annotation], resource: SlicedVolumeResource, seg_slice_axis: int, ) -> np.ndarray | None: diff --git a/datamint/dataset/sliced_video_dataset.py b/datamint/dataset/sliced_video_dataset.py index 7621b8bb..abe1e02e 100644 --- a/datamint/dataset/sliced_video_dataset.py +++ b/datamint/dataset/sliced_video_dataset.py @@ -5,21 +5,23 @@ enabling training of 2D models on temporal medical imaging data. """ from __future__ import annotations + import hashlib import logging -from typing import Any, TYPE_CHECKING -from typing_extensions import override from collections.abc import Sequence +from typing import TYPE_CHECKING, Any +import albumentations import numpy as np import torch from torch import Tensor -import albumentations +from typing_extensions import override -from .base import DatamintBaseDataset -from .annotation_processor import AnnotationProcessor from datamint.entities.cache_manager import CacheManager +from .annotation_processor import AnnotationProcessor +from .base import DatamintBaseDataset + if TYPE_CHECKING: from datamint.entities import Annotation from datamint.entities.sliced_video_resource import SlicedVideoResource @@ -145,9 +147,9 @@ def from_dataset( def _expand_resources( self, resources: Sequence, - resource_annotations: Sequence[Sequence['Annotation']], + resource_annotations: Sequence[Sequence[Annotation]], frame_cache: CacheManager, - ) -> tuple[list['SlicedVideoResource'], list[Sequence['Annotation']]]: + ) -> tuple[list[SlicedVideoResource], list[Sequence[Annotation]]]: """Expand video resources into per-frame proxy resources. Args: @@ -173,8 +175,8 @@ def _expand_resources( def _load_frame_segmentations( self, - annotations: Sequence['Annotation'], - resource: 'SlicedVideoResource', + annotations: Sequence[Annotation], + resource: SlicedVideoResource, ) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray], dict[str, list]]: """Load segmentations for a specific video frame, with caching. @@ -195,9 +197,9 @@ def _load_frame_segmentations( image_seg_anns = [a for a in seg_anns if a.scope == 'image'] frame_seg_anns = [a for a in seg_anns if a.scope == 'frame'] - uniq_authors = set( + uniq_authors = { self.annotation_processor.get_author(a) for a in seg_anns - ) + } segmentations: dict[str, list[np.ndarray]] = {a: [] for a in uniq_authors} seg_labels: dict[str, list[int]] = {a: [] for a in uniq_authors} seg_metainfos: dict[str, list] = {a: [] for a in uniq_authors} @@ -241,8 +243,8 @@ def _load_frame_segmentations( def _fetch_frame_seg_annotation( self, - ann: 'Annotation', - resource: 'SlicedVideoResource', + ann: Annotation, + resource: SlicedVideoResource, ) -> np.ndarray: """Load an image-scoped segmentation and extract the frame, with caching. @@ -275,8 +277,8 @@ def _fetch_frame_seg_annotation( def _fetch_frame_seg_group( self, - fr_anns: list['Annotation'], - resource: 'SlicedVideoResource', + fr_anns: list[Annotation], + resource: SlicedVideoResource, ) -> np.ndarray | None: """Collate frame-level segmentation annotations and extract the frame. diff --git a/datamint/dataset/split_result.py b/datamint/dataset/split_result.py index 42756e18..adcc7e0d 100644 --- a/datamint/dataset/split_result.py +++ b/datamint/dataset/split_result.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING if TYPE_CHECKING: from .base import DatamintBaseDataset diff --git a/datamint/entities/__init__.py b/datamint/entities/__init__.py index ee266fc6..d2ade266 100644 --- a/datamint/entities/__init__.py +++ b/datamint/entities/__init__.py @@ -1,21 +1,27 @@ """DataMint entities package.""" from .annotations.annotation import Annotation +from .annotations.annotation_spec import AnnotationSpec from .base_entity import BaseEntity, BaseEntityModel +from .cache_manager import CacheManager from .channel import Channel, ChannelResourceData -from .project import Project -from .resource import Resource -from .user import User # new export from .datasetinfo import DatasetInfo -from .cache_manager import CacheManager from .inferencejob import InferenceJob -from .annotations.annotation_spec import AnnotationSpec +from .project import Project from .project_resource_split import ProjectResourceSplit -from .resource import LocalResource -from .resources import DICOMResource, ImageResource, NiftiResource, VideoResource, VolumeResource +from .resource import LocalResource, Resource +from .resources import ( + DICOMResource, + ImageResource, + NiftiResource, + VideoResource, + VolumeResource, +) +from .user import User # new export __all__ = [ 'Annotation', + 'AnnotationSpec', 'BaseEntity', 'BaseEntityModel', 'CacheManager', @@ -30,8 +36,7 @@ 'Project', 'ProjectResourceSplit', 'Resource', - 'VideoResource', - 'VolumeResource', 'User', - 'AnnotationSpec' + 'VideoResource', + 'VolumeResource' ] diff --git a/datamint/entities/annotation_worklist.py b/datamint/entities/annotation_worklist.py index 1e7194cc..b8505671 100644 --- a/datamint/entities/annotation_worklist.py +++ b/datamint/entities/annotation_worklist.py @@ -5,7 +5,7 @@ from datamint.entities.annotations.annotation_spec import AnnotationSpec -from .base_entity import BaseEntity, MISSING_FIELD +from .base_entity import MISSING_FIELD, BaseEntity _LOGGER = logging.getLogger(__name__) diff --git a/datamint/entities/annotations/__init__.py b/datamint/entities/annotations/__init__.py index fcb3c3a2..31e39329 100644 --- a/datamint/entities/annotations/__init__.py +++ b/datamint/entities/annotations/__init__.py @@ -1,12 +1,12 @@ -from .image_classification import ImageClassification -from .image_segmentation import ImageSegmentation from .annotation import Annotation, _normalize_annotation_data from .box_annotation import BoxAnnotation from .geometry import BoxGeometry, CoordinateSystem, Geometry, LineGeometry +from .image_classification import ImageClassification +from .image_segmentation import ImageSegmentation from .line_annotation import LineAnnotation from .numeric_annotation import NumericAnnotation -from .volume_segmentation import VolumeSegmentation from .types import AnnotationType +from .volume_segmentation import VolumeSegmentation def annotation_from_dict(data: dict) -> Annotation: @@ -58,17 +58,17 @@ def annotation_from_dict(data: dict) -> Annotation: __all__ = [ + "Annotation", + "AnnotationType", "BoxAnnotation", "BoxGeometry", - "ImageClassification", - "ImageSegmentation", - "Annotation", "CoordinateSystem", "Geometry", + "ImageClassification", + "ImageSegmentation", "LineAnnotation", "LineGeometry", "NumericAnnotation", "VolumeSegmentation", - "AnnotationType", "annotation_from_dict", ] diff --git a/datamint/entities/annotations/annotation.py b/datamint/entities/annotations/annotation.py index 0d46ded6..e19de4b8 100644 --- a/datamint/entities/annotations/annotation.py +++ b/datamint/entities/annotations/annotation.py @@ -5,20 +5,21 @@ records returned by the DataMint API. """ -from datetime import datetime import logging +from datetime import datetime from typing import TYPE_CHECKING, Any, Literal, overload from pydantic import ConfigDict, Field, PrivateAttr, field_validator from datamint.types import CacheMode, ImagingData -from ..base_entity import BaseEntity, MISSING_FIELD +from ..base_entity import MISSING_FIELD, BaseEntity from ..cache_manager import CacheManager from .types import AnnotationType if TYPE_CHECKING: from datamint.api.endpoints.annotations_api import AnnotationsApi + from ..resource import Resource @@ -178,7 +179,7 @@ class Annotation(AnnotationBase): def __init__(self, **data): """Initialize the annotation entity.""" super().__init__(**data) - self._resource: 'Resource | None' = None + self._resource: Resource | None = None @property def _cache(self) -> CacheManager[bytes]: @@ -323,7 +324,7 @@ def from_dict(cls, data: dict[str, Any]) -> 'Annotation': raise ValueError(f"Segmentation annotations must have an associated file. {data}") # Create instance with only valid fields - valid_fields = {f for f in cls.model_fields.keys()} + valid_fields = {f for f in cls.model_fields} filtered_data = {k: v for k, v in converted_data.items() if k in valid_fields} return cls(**filtered_data) diff --git a/datamint/entities/annotations/annotation_spec.py b/datamint/entities/annotations/annotation_spec.py index 075bc70b..2a81a0a1 100644 --- a/datamint/entities/annotations/annotation_spec.py +++ b/datamint/entities/annotations/annotation_spec.py @@ -1,6 +1,7 @@ from typing import Any from pydantic import BaseModel, ConfigDict + from .types import AnnotationType diff --git a/datamint/entities/annotations/base_segmentation.py b/datamint/entities/annotations/base_segmentation.py index ff08d779..464c17c2 100644 --- a/datamint/entities/annotations/base_segmentation.py +++ b/datamint/entities/annotations/base_segmentation.py @@ -14,12 +14,13 @@ from typing import Annotated, Any, Literal, overload import numpy as np +from nibabel.nifti1 import Nifti1Image from PIL import Image -from pydantic import BeforeValidator, PlainSerializer, Field +from pydantic import BeforeValidator, Field, PlainSerializer -from .annotation import Annotation from datamint.types import CacheMode, ImagingData -from nibabel.nifti1 import Nifti1Image + +from .annotation import Annotation _LOGGER = logging.getLogger(__name__) diff --git a/datamint/entities/annotations/box_annotation.py b/datamint/entities/annotations/box_annotation.py index e83c4b3a..53faee6c 100644 --- a/datamint/entities/annotations/box_annotation.py +++ b/datamint/entities/annotations/box_annotation.py @@ -1,11 +1,10 @@ from __future__ import annotations -from pathlib import Path from typing import Any import pydicom -from pydantic import field_validator from nibabel.nifti1 import Nifti1Image +from pydantic import field_validator from .base_geometry import BaseGeometryAnnotation from .geometry import BoxGeometry, CoordinateSystem @@ -37,7 +36,7 @@ def from_points( metadata: pydicom.Dataset | Nifti1Image | None = None, coords_system: CoordinateSystem = 'pixel', **kwargs: Any, - ) -> 'BoxAnnotation': + ) -> BoxAnnotation: geometry = BoxGeometry.from_coordinates( point1, point2, diff --git a/datamint/entities/annotations/geometry.py b/datamint/entities/annotations/geometry.py index 9fca0975..ebdc5303 100644 --- a/datamint/entities/annotations/geometry.py +++ b/datamint/entities/annotations/geometry.py @@ -1,11 +1,13 @@ from __future__ import annotations -from typing import Any, ClassVar, Literal, TypeAlias + import logging -from nibabel.nifti1 import Nifti1Image +from typing import Any, ClassVar, Literal, TypeAlias + import numpy as np -from medimgkit.dicom_utils import get_slice_orientation import pydicom from medimgkit import ViewPlane, dicom_utils, nifti_utils +from medimgkit.dicom_utils import get_slice_orientation +from nibabel.nifti1 import Nifti1Image from pydantic import BaseModel, ConfigDict, field_validator CoordinateSystem: TypeAlias = Literal['pixel', 'patient'] @@ -137,7 +139,7 @@ def from_coordinates( slice_plane: ViewPlane | None = None, frame_index: int | None = None, metadata: pydicom.Dataset | Nifti1Image | None = None, - ) -> '_TwoPointGeometry': + ) -> _TwoPointGeometry: if coords_system == 'pixel': return cls._from_pixel_coordinates( point1, @@ -236,7 +238,7 @@ def _from_pixel_coordinates( frame_index: int | None = None, slice_plane: ViewPlane | None = None, metadata: pydicom.Dataset | Nifti1Image | None = None, - ) -> '_TwoPointGeometry': + ) -> _TwoPointGeometry: point1 = _normalize_point(point1) point2 = _normalize_point(point2) @@ -308,7 +310,7 @@ def from_coordinates( slice_plane: ViewPlane | None = None, frame_index: int | None = None, metadata: pydicom.Dataset | Nifti1Image | None = None, - ) -> 'BoxGeometry': + ) -> BoxGeometry: if coords_system == 'pixel': return cls._from_pixel_coordinates( point1, @@ -392,7 +394,7 @@ def _from_pixel_coordinates( frame_index: int | None = None, slice_plane: ViewPlane | None = None, metadata: pydicom.Dataset | Nifti1Image | None = None, - ) -> 'BoxGeometry': + ) -> BoxGeometry: normalized_point1 = _normalize_point(point1) normalized_point2 = _normalize_point(point2) diff --git a/datamint/entities/annotations/image_segmentation.py b/datamint/entities/annotations/image_segmentation.py index 7e065002..3faffc22 100644 --- a/datamint/entities/annotations/image_segmentation.py +++ b/datamint/entities/annotations/image_segmentation.py @@ -11,8 +11,8 @@ import numpy as np from PIL import Image -from .types import AnnotationType from .base_segmentation import BaseSegmentationAnnotation +from .types import AnnotationType _LOGGER = logging.getLogger(__name__) diff --git a/datamint/entities/annotations/line_annotation.py b/datamint/entities/annotations/line_annotation.py index f249f675..a1c49d92 100644 --- a/datamint/entities/annotations/line_annotation.py +++ b/datamint/entities/annotations/line_annotation.py @@ -1,10 +1,10 @@ from __future__ import annotations -from medimgkit import ViewPlane from typing import Any -from nibabel.nifti1 import Nifti1Image import pydicom +from medimgkit import ViewPlane +from nibabel.nifti1 import Nifti1Image from pydantic import field_validator from .base_geometry import BaseGeometryAnnotation @@ -37,7 +37,7 @@ def from_points( metadata: pydicom.Dataset | Nifti1Image | None = None, coords_system: CoordinateSystem = 'pixel', **kwargs: Any, - ) -> 'LineAnnotation': + ) -> LineAnnotation: geometry = LineGeometry.from_coordinates( point1, point2, diff --git a/datamint/entities/annotations/numeric_annotation.py b/datamint/entities/annotations/numeric_annotation.py index e093bfff..405debba 100644 --- a/datamint/entities/annotations/numeric_annotation.py +++ b/datamint/entities/annotations/numeric_annotation.py @@ -10,7 +10,7 @@ class NumericAnnotation(Annotation): def __init__( self, name: str | None = None, - value: int | float | None = None, + value: float | None = None, units: str | None = None, confiability: float = 1.0, **kwargs: Any, diff --git a/datamint/entities/annotations/volume_segmentation.py b/datamint/entities/annotations/volume_segmentation.py index 1af4dfe7..993924df 100644 --- a/datamint/entities/annotations/volume_segmentation.py +++ b/datamint/entities/annotations/volume_segmentation.py @@ -4,11 +4,13 @@ annotations in medical imaging volumes. """ -from .base_segmentation import BaseSegmentationAnnotation -from .types import AnnotationType +import logging + import numpy as np from nibabel.nifti1 import Nifti1Image -import logging + +from .base_segmentation import BaseSegmentationAnnotation +from .types import AnnotationType _LOGGER = logging.getLogger(__name__) @@ -92,9 +94,8 @@ def from_semantic_segmentation(cls, ... class_map={1: 'tumor'}, # or just ``class_map='tumor'`` ... ) """ - if not kwargs.get('identifier'): - if isinstance(class_map, str): - kwargs['identifier'] = class_map + if not kwargs.get('identifier') and isinstance(class_map, str): + kwargs['identifier'] = class_map # Step 1: Convert Nifti1Image to numpy if needed if isinstance(segmentation, Nifti1Image): diff --git a/datamint/entities/base_entity.py b/datamint/entities/base_entity.py index 5282347a..f6fdfdb6 100644 --- a/datamint/entities/base_entity.py +++ b/datamint/entities/base_entity.py @@ -1,12 +1,11 @@ import logging import sys -from html import escape -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any from pydantic import BaseModel, ConfigDict, PrivateAttr +from datamint._repr_utils import render_html_card, render_text_block from datamint.types import CacheMode -from datamint._repr_utils import render_text_block, render_html_card if TYPE_CHECKING: from datamint.api.entity_base_api import EntityBaseApi @@ -74,8 +73,8 @@ def __str__(self) -> str: def __init__(self, **data): super().__init__(**data) - for field_name in self.__pydantic_fields__.keys(): - if hasattr(self, field_name) and type(getattr(self, field_name)) == str and getattr(self, field_name) == MISSING_FIELD: + for field_name in self.__pydantic_fields__: + if hasattr(self, field_name) and isinstance(getattr(self, field_name), str) and getattr(self, field_name) == MISSING_FIELD: delattr(self, field_name) def asdict(self) -> dict[str, Any]: @@ -93,7 +92,7 @@ def model_post_init(self, __context: Any) -> None: class_name = self.__class__.__name__ have_to_log = False - for key in self.__pydantic_extra__.keys(): + for key in self.__pydantic_extra__: warning_key = (class_name, key) if warning_key not in _LOGGED_WARNINGS: @@ -105,7 +104,7 @@ def model_post_init(self, __context: Any) -> None: def is_attr_missing(self, attr_name: str) -> bool: """Check if a value is the MISSING_FIELD sentinel.""" - if attr_name not in self.__pydantic_fields__.keys(): + if attr_name not in self.__pydantic_fields__: raise AttributeError(f"Attribute '{attr_name}' not found in entity of type '{self.__class__.__name__}'") if not hasattr(self, attr_name): return True @@ -117,7 +116,7 @@ def has_missing_attrs(self) -> bool: Returns: True if any attribute is MISSING_FIELD, False otherwise """ - return any(self.is_attr_missing(attr_name) for attr_name in self.__pydantic_fields__.keys()) + return any(self.is_attr_missing(attr_name) for attr_name in self.__pydantic_fields__) class BaseEntity(BaseEntityModel): @@ -166,7 +165,7 @@ def _ensure_attr(self, attr_name: str) -> None: Args: attr_name: Name of the attribute to check and ensure """ - if attr_name not in self.__pydantic_fields__.keys(): + if attr_name not in self.__pydantic_fields__: raise AttributeError(f"Attribute '{attr_name}' not found in entity of type '{self.__class__.__name__}'") if self.is_attr_missing(attr_name): diff --git a/datamint/entities/cache_manager.py b/datamint/entities/cache_manager.py index ed093380..80c81f5a 100644 --- a/datamint/entities/cache_manager.py +++ b/datamint/entities/cache_manager.py @@ -4,18 +4,20 @@ with automatic validation against server versions to ensure data freshness. """ +import gzip import hashlib import json import logging import pickle -import gzip -import numpy as np from collections.abc import Iterator from datetime import datetime from pathlib import Path -from typing import Any, TypeVar, Generic -from pydantic import BaseModel +from typing import Any, Generic, TypeVar + +import numpy as np from cachetools import LRUCache +from pydantic import BaseModel + # import appdirs import datamint.configs @@ -125,7 +127,7 @@ def invalidate_memory(self, entity_id: str, data_key: str | None = None) -> None if self._memory_cache is None: return keys_to_remove = [] - for eid, dkey, vhash in self._memory_cache.keys(): + for eid, dkey, vhash in self._memory_cache: if eid != entity_id: continue if data_key is not None and dkey != data_key: @@ -259,7 +261,7 @@ def get( if mem_data is not None: return mem_data - cached_metadata, data_path = self._get_validated_metadata(entity_id, data_key, version_info) + cached_metadata, _data_path = self._get_validated_metadata(entity_id, data_key, version_info) if cached_metadata is None: return None @@ -288,7 +290,7 @@ def get_path( Returns: Path to cached data if valid, None if cache miss or invalid """ - cached_metadata, data_path = self._get_validated_metadata(entity_id, data_key, version_info) + _cached_metadata, data_path = self._get_validated_metadata(entity_id, data_key, version_info) return data_path def get_expected_path(self, entity_id: str, data_key: str) -> Path: diff --git a/datamint/entities/channel.py b/datamint/entities/channel.py index 9b6637bb..51df3fc9 100644 --- a/datamint/entities/channel.py +++ b/datamint/entities/channel.py @@ -1,5 +1,6 @@ -from pydantic import ConfigDict, BaseModel -from datetime import datetime + +from pydantic import BaseModel, ConfigDict + from datamint.entities.base_entity import BaseEntity diff --git a/datamint/entities/datasetinfo.py b/datamint/entities/datasetinfo.py index aacfa120..6c2ab2f5 100644 --- a/datamint/entities/datasetinfo.py +++ b/datamint/entities/datasetinfo.py @@ -1,15 +1,13 @@ """Dataset entity module for DataMint API.""" import logging -from typing import TYPE_CHECKING, Sequence +from typing import TYPE_CHECKING + from pydantic import PrivateAttr -from .base_entity import BaseEntity, MISSING_FIELD +from .base_entity import BaseEntity if TYPE_CHECKING: - from datamint.api.client import Api - from .resource import Resource - from .project import Project from datamint.api.endpoints.datasetsinfo_api import DatasetsInfoApi logger = logging.getLogger(__name__) diff --git a/datamint/entities/deployjob.py b/datamint/entities/deployjob.py index 1abc8c91..94405697 100644 --- a/datamint/entities/deployjob.py +++ b/datamint/entities/deployjob.py @@ -1,6 +1,5 @@ -from collections.abc import Callable import logging -from typing import Any +from collections.abc import Callable from pydantic import field_validator diff --git a/datamint/entities/inferencejob.py b/datamint/entities/inferencejob.py index d778fad5..f297d8db 100644 --- a/datamint/entities/inferencejob.py +++ b/datamint/entities/inferencejob.py @@ -1,11 +1,11 @@ from __future__ import annotations -from collections.abc import Callable import logging -from typing import Any, TYPE_CHECKING +from collections.abc import Callable +from typing import TYPE_CHECKING, Any -from datamint.entities.base_entity import BaseEntity, MISSING_FIELD from datamint.entities.annotations import annotation_from_dict +from datamint.entities.base_entity import MISSING_FIELD, BaseEntity if TYPE_CHECKING: from datamint.api.endpoints.inference_api import InferenceApi @@ -186,7 +186,7 @@ def is_finished(self) -> bool: return self.status.lower() in {'completed', 'failed', 'cancelled', 'error'} @property - def predictions(self) -> 'list[list[Annotation]] | None': + def predictions(self) -> list[list[Annotation]] | None: """ Returns a list of annotations resulting from this inference job, if available. diff --git a/datamint/entities/project.py b/datamint/entities/project.py index 0bc5870f..edcfe326 100644 --- a/datamint/entities/project.py +++ b/datamint/entities/project.py @@ -1,20 +1,19 @@ """Project entity module for DataMint API.""" import logging -from typing import Literal, TYPE_CHECKING +import webbrowser from collections.abc import Sequence +from typing import TYPE_CHECKING, Any, Literal -from .base_entity import BaseEntity, MISSING_FIELD -from typing import Any -import webbrowser -from pydantic import PrivateAttr, Field, BaseModel -from functools import cached_property +from pydantic import BaseModel, Field, PrivateAttr + +from .base_entity import MISSING_FIELD, BaseEntity if TYPE_CHECKING: from datamint.api.endpoints.projects_api import ProjectsApi - from .resource import Resource - from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.annotation_worklist import AnnotationWorklist + from .resource import Resource + logger = logging.getLogger(__name__) diff --git a/datamint/entities/resource.py b/datamint/entities/resource.py index 13488054..7e8b6df1 100644 --- a/datamint/entities/resource.py +++ b/datamint/entities/resource.py @@ -1,31 +1,34 @@ """Resource entity module for DataMint API.""" -from collections.abc import Sequence -from abc import ABC, abstractmethod -from datetime import datetime import logging -from pathlib import Path -from typing import TYPE_CHECKING, Any, ClassVar, Literal, overload -from typing_extensions import override import urllib.parse import urllib.request import webbrowser +from abc import ABC, abstractmethod +from collections.abc import Sequence +from datetime import datetime +from pathlib import Path +from typing import TYPE_CHECKING, Any, ClassVar, Literal, overload -from pydantic import PrivateAttr +from medimgkit.nifti_utils import NIFTI_MIMES +from pydantic import Field, PrivateAttr +from typing_extensions import override -from .base_entity import BaseEntity, MISSING_FIELD -from .cache_manager import CacheManager from datamint.api.base_api import BaseApi from datamint.types import CacheMode -from medimgkit.nifti_utils import NIFTI_MIMES + +from .base_entity import MISSING_FIELD, BaseEntity +from .cache_manager import CacheManager if TYPE_CHECKING: - from datamint.api.endpoints.resources_api import ResourcesApi + import numpy as np from medimgkit import ViewPlane - from .annotations.annotation import Annotation - from .annotations import AnnotationType + + from datamint.api.endpoints.resources_api import ResourcesApi from datamint.types import ImagingData - import numpy as np + + from .annotations import AnnotationType + from .annotations.annotation import Annotation from .sliced_resource import SlicedVolumeResource from .sliced_video_resource import SlicedVideoResource @@ -114,7 +117,6 @@ def fetch_file_data( Returns: File data (format depends on auto_convert and file type) """ - pass @property def size_mb(self) -> float: @@ -138,7 +140,7 @@ def is_image(self) -> bool: """Check if the resource is a single-frame image.""" if not self.mimetype: return False - return self.mimetype.startswith('image/') and not self.mimetype == 'image/nifti' + return self.mimetype.startswith('image/') and self.mimetype != 'image/nifti' def is_dicom(self) -> bool: """Check if the resource is a DICOM file. @@ -233,7 +235,7 @@ class Resource(BaseResource): published: bool deleted: bool upload_mechanism: str | None = None - metadata: dict[str, Any] = {} + metadata: dict[str, Any] = Field(default_factory=dict) source_filepath: str | None = None # projects: list[dict[str, Any]] | None = None published_on: str | None = None @@ -326,13 +328,7 @@ def matches_payload( for prefix in cls.mimetype_prefixes ): return True - if filename_norm and any( - filename_norm.endswith(suffix.casefold()) - for suffix in cls.filename_suffixes - ): - return True - - return False + return bool(filename_norm and any(filename_norm.endswith(suffix.casefold()) for suffix in cls.filename_suffixes)) @classmethod def _infer_specialized_resource_class(cls, **kwargs) -> type['Resource']: @@ -750,7 +746,7 @@ def __init__(self, raw_data: Raw bytes of the file data convert_to_bytes: If True and local_filepath is provided, read file into raw_data """ - from medimgkit.format_detection import guess_type, DEFAULT_MIME_TYPE + from medimgkit.format_detection import DEFAULT_MIME_TYPE, guess_type from medimgkit.modality_detector import detect_modality if raw_data is None and local_filepath is None: @@ -920,7 +916,7 @@ def fetch_file_data( if auto_convert: try: - mimetype, ext = BaseApi._determine_mimetype(img_data, self.mimetype) + mimetype, _ext = BaseApi._determine_mimetype(img_data, self.mimetype) img_data = BaseApi.convert_format(img_data, mimetype=mimetype, file_path=local_filepath) diff --git a/datamint/entities/resources/nifti_resource.py b/datamint/entities/resources/nifti_resource.py index fc657e4e..c07715c2 100644 --- a/datamint/entities/resources/nifti_resource.py +++ b/datamint/entities/resources/nifti_resource.py @@ -2,9 +2,10 @@ from typing import ClassVar -from .volume_resource import VolumeResource from medimgkit.nifti_utils import NIFTI_MIMES +from .volume_resource import VolumeResource + class NiftiResource(VolumeResource): """Represents a NIfTI volume resource.""" diff --git a/datamint/entities/resources/video_resource.py b/datamint/entities/resources/video_resource.py index d369c97c..2884f45a 100644 --- a/datamint/entities/resources/video_resource.py +++ b/datamint/entities/resources/video_resource.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, ClassVar, TYPE_CHECKING +from typing import TYPE_CHECKING, Any, ClassVar from ..resource import Resource @@ -64,5 +64,5 @@ def get_depth(self) -> int: return frame_count - def iter_frames(self) -> list['SlicedVideoResource']: + def iter_frames(self) -> list[SlicedVideoResource]: return super().iter_frames() \ No newline at end of file diff --git a/datamint/entities/resources/volume_resource.py b/datamint/entities/resources/volume_resource.py index fc116421..02694d50 100644 --- a/datamint/entities/resources/volume_resource.py +++ b/datamint/entities/resources/volume_resource.py @@ -1,12 +1,13 @@ from __future__ import annotations -from typing import ClassVar, TYPE_CHECKING +from typing import TYPE_CHECKING, ClassVar from ..resource import Resource if TYPE_CHECKING: - from medimgkit import ViewPlane import numpy as np + from medimgkit import ViewPlane + from ..sliced_resource import SlicedVolumeResource @@ -27,11 +28,11 @@ def get_depth(self) -> int: raise ValueError(f"Cannot determine frame count for volume resource {self.filename!r}") return frame_count - def get_slice_resource(self, axis: 'ViewPlane', index: int) -> 'SlicedVolumeResource': + def get_slice_resource(self, axis: ViewPlane, index: int) -> SlicedVolumeResource: return super().get_slice_resource(axis, index) - def get_slice(self, axis: 'ViewPlane', index: int) -> 'np.ndarray': + def get_slice(self, axis: ViewPlane, index: int) -> np.ndarray: return super().get_slice(axis, index) - def iter_slices(self, axis: 'ViewPlane') -> list['SlicedVolumeResource']: + def iter_slices(self, axis: ViewPlane) -> list[SlicedVolumeResource]: return super().iter_slices(axis) \ No newline at end of file diff --git a/datamint/entities/sliced_resource.py b/datamint/entities/sliced_resource.py index ff8fa48d..d6e0b2ab 100644 --- a/datamint/entities/sliced_resource.py +++ b/datamint/entities/sliced_resource.py @@ -1,12 +1,16 @@ from __future__ import annotations + import gzip import logging -from typing import Any, TYPE_CHECKING from functools import cached_property +from typing import TYPE_CHECKING, Any + +import numpy as np +from medimgkit import ViewPlane, dicom_utils, nifti_utils from medimgkit.readers import read_array_normalized -from medimgkit import dicom_utils, nifti_utils, ViewPlane + from datamint.entities.cache_manager import CacheManager -import numpy as np + from .sliced_resource_base import SlicedResourceBase if TYPE_CHECKING: diff --git a/datamint/entities/sliced_resource_base.py b/datamint/entities/sliced_resource_base.py index f5200f8c..3f1c6e3f 100644 --- a/datamint/entities/sliced_resource_base.py +++ b/datamint/entities/sliced_resource_base.py @@ -1,7 +1,9 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any + import logging from functools import cached_property +from typing import TYPE_CHECKING, Any + from medimgkit.readers import read_array_normalized if TYPE_CHECKING: diff --git a/datamint/entities/sliced_video_resource.py b/datamint/entities/sliced_video_resource.py index ce243dac..6d5a3bea 100644 --- a/datamint/entities/sliced_video_resource.py +++ b/datamint/entities/sliced_video_resource.py @@ -5,13 +5,16 @@ data — videos always slice along the frame (temporal) axis. """ from __future__ import annotations + import gzip import logging -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any +import numpy as np from medimgkit.readers import read_array_normalized + from datamint.entities.cache_manager import CacheManager -import numpy as np + from .sliced_resource_base import SlicedResourceBase if TYPE_CHECKING: diff --git a/datamint/examples/__init__.py b/datamint/examples/__init__.py index 7d17ae06..d353f357 100644 --- a/datamint/examples/__init__.py +++ b/datamint/examples/__init__.py @@ -1,2 +1,5 @@ -from .example_projects import ProjectMR -from . import bccd_dataset, busi_dataset, synapse_dataset, fracatlas_dataset \ No newline at end of file +from . import bccd_dataset as bccd_dataset +from . import busi_dataset as busi_dataset +from . import fracatlas_dataset as fracatlas_dataset +from . import synapse_dataset as synapse_dataset +from .example_projects import ProjectMR as ProjectMR diff --git a/datamint/examples/example_projects.py b/datamint/examples/example_projects.py index b4adf90b..806b6994 100644 --- a/datamint/examples/example_projects.py +++ b/datamint/examples/example_projects.py @@ -1,12 +1,14 @@ -import requests import io -from datamint import Api import logging -from PIL import Image + import numpy as np -from datamint.entities import Project, Resource +import requests +from PIL import Image from pydicom.data import get_testdata_file +from datamint import Api +from datamint.entities import Project, Resource + _LOGGER = logging.getLogger(__name__) diff --git a/datamint/examples/synapse_dataset.py b/datamint/examples/synapse_dataset.py index 3b34825a..7a6332ee 100644 --- a/datamint/examples/synapse_dataset.py +++ b/datamint/examples/synapse_dataset.py @@ -1,7 +1,7 @@ import logging -import numpy as np import nibabel as nib +import numpy as np from tqdm.auto import tqdm from datamint import Api diff --git a/datamint/exceptions.py b/datamint/exceptions.py index 8482d776..f8c038b4 100644 --- a/datamint/exceptions.py +++ b/datamint/exceptions.py @@ -1,6 +1,5 @@ class DatamintException(Exception): """Base class for all Datamint exceptions.""" - pass # --------------------------------------------------------------------------- @@ -9,12 +8,10 @@ class DatamintException(Exception): class AuthenticationError(DatamintException): """Raised when the API key is missing or rejected (HTTP 401).""" - pass class PermissionDeniedError(DatamintException): """Raised when the authenticated user lacks permission for the requested operation (HTTP 403).""" - pass # --------------------------------------------------------------------------- @@ -89,7 +86,6 @@ def __str__(self) -> str: class ValidationError(DatamintException): """Raised when the server rejects a request due to invalid input (HTTP 400/422).""" - pass # --------------------------------------------------------------------------- @@ -98,7 +94,6 @@ class ValidationError(DatamintException): class NetworkError(DatamintException): """Raised on connection failures, SSL errors, or other transport-level problems.""" - pass # --------------------------------------------------------------------------- @@ -169,4 +164,3 @@ class JobTimeoutError(DatamintException, TimeoutError): Subclasses both DatamintException and the built-in TimeoutError so callers catching either one will handle it correctly. """ - pass diff --git a/datamint/lightning/__init__.py b/datamint/lightning/__init__.py index acd15881..d1762ab5 100644 --- a/datamint/lightning/__init__.py +++ b/datamint/lightning/__init__.py @@ -4,34 +4,34 @@ from .trainers import ( BaseTrainer, ClassificationTrainer, + DeepLabV3PlusTrainer, + EfficientNetV2Trainer, ImageClassificationTrainer, + NNUNetTrainer, + SegmentationTrainer, SemanticSegmentation2DTrainer, SemanticSegmentation3DTrainer, - SegmentationTrainer, - VolumeSegmentationTrainer, - UNetPPTrainer, - DeepLabV3PlusTrainer, TransUNetTrainer, - YOLOXTrainer, - NNUNetTrainer, + UNetPPTrainer, UNETRPPTrainer, - EfficientNetV2Trainer, + VolumeSegmentationTrainer, + YOLOXTrainer, ) __all__ = [ - "DatamintDataModule", "BaseTrainer", "ClassificationTrainer", + "DatamintDataModule", + "DeepLabV3PlusTrainer", + "EfficientNetV2Trainer", "ImageClassificationTrainer", + "NNUNetTrainer", + "SegmentationTrainer", "SemanticSegmentation2DTrainer", "SemanticSegmentation3DTrainer", - "SegmentationTrainer", - "VolumeSegmentationTrainer", - "UNetPPTrainer", - "DeepLabV3PlusTrainer", "TransUNetTrainer", - "YOLOXTrainer", - "NNUNetTrainer", "UNETRPPTrainer", - "EfficientNetV2Trainer", + "UNetPPTrainer", + "VolumeSegmentationTrainer", + "YOLOXTrainer", ] diff --git a/datamint/lightning/datamodule.py b/datamint/lightning/datamodule.py index 27f5df73..fbc3cf36 100644 --- a/datamint/lightning/datamodule.py +++ b/datamint/lightning/datamodule.py @@ -256,7 +256,7 @@ def test_dataloader(self) -> DataLoader: collate_fn=self.collate_fn if self.collate_fn is not None else self.dataset.get_collate_fn(), ) - def get_mlflow_dataset_split(self, split: str) -> 'DatamintMLflowDataset | None': + def get_mlflow_dataset_split(self, split: str) -> DatamintMLflowDataset | None: """Return a :class:`~datamint.mlflow.data.DatamintMLflowDataset` for the given split. Delegates to the corresponding split dataset's @@ -300,7 +300,7 @@ def has_val_split(self) -> bool: val = parts.get("val") return val is not None and len(val) > 0 - def get_dataset_split(self, split: str) -> 'DatamintBaseDataset | None': + def get_dataset_split(self, split: str) -> DatamintBaseDataset | None: """Return the Datamint dataset for the given split. Falls back to the full dataset when the requested split is not available.""" split_ds_map = { 'train': self._train_dataset, diff --git a/datamint/lightning/trainers/__init__.py b/datamint/lightning/trainers/__init__.py index 040615c9..84ef1a59 100644 --- a/datamint/lightning/trainers/__init__.py +++ b/datamint/lightning/trainers/__init__.py @@ -1,28 +1,36 @@ """Specialized trainers for end-to-end Datamint workflows.""" from .base_trainer import BaseTrainer -from .segmentation_trainer import SegmentationTrainer -from .seg2d_trainer import SemanticSegmentation2DTrainer -from .seg3d_trainer import SemanticSegmentation3DTrainer from .classification_trainer import ClassificationTrainer, ImageClassificationTrainer from .detection_trainer import DetectionTrainer -from .specialized import UNetPPTrainer, DeepLabV3PlusTrainer, TransUNetTrainer, UNETRPPTrainer, NNUNetTrainer, YOLOXTrainer, EfficientNetV2Trainer +from .seg2d_trainer import SemanticSegmentation2DTrainer +from .seg3d_trainer import SemanticSegmentation3DTrainer +from .segmentation_trainer import SegmentationTrainer +from .specialized import ( + DeepLabV3PlusTrainer, + EfficientNetV2Trainer, + NNUNetTrainer, + TransUNetTrainer, + UNetPPTrainer, + UNETRPPTrainer, + YOLOXTrainer, +) from .vol_seg_trainer import VolumeSegmentationTrainer __all__ = [ "BaseTrainer", + "ClassificationTrainer", + "DeepLabV3PlusTrainer", + "DetectionTrainer", + "EfficientNetV2Trainer", + "ImageClassificationTrainer", + "NNUNetTrainer", "SegmentationTrainer", "SemanticSegmentation2DTrainer", "SemanticSegmentation3DTrainer", - "VolumeSegmentationTrainer", - "UNetPPTrainer", - "DeepLabV3PlusTrainer", "TransUNetTrainer", "UNETRPPTrainer", - "ClassificationTrainer", - "ImageClassificationTrainer", - "NNUNetTrainer", - "YOLOXTrainer", - "EfficientNetV2Trainer", - "DetectionTrainer" + "UNetPPTrainer", + "VolumeSegmentationTrainer", + "YOLOXTrainer" ] diff --git a/datamint/lightning/trainers/base_trainer.py b/datamint/lightning/trainers/base_trainer.py index c0d7e6ea..8ca97c3a 100644 --- a/datamint/lightning/trainers/base_trainer.py +++ b/datamint/lightning/trainers/base_trainer.py @@ -7,23 +7,24 @@ import logging from abc import ABC, abstractmethod -from functools import cached_property from collections.abc import Callable, Mapping -from typing import Any, TYPE_CHECKING, cast +from functools import cached_property +from typing import TYPE_CHECKING, Any, ClassVar import lightning as L -from torch import nn import mlflow +from torch import nn +from datamint._repr_utils import render_html_card, render_text_block from datamint.dataset.base import DatamintBaseDataset from datamint.lightning.datamodule import DatamintDataModule +from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule from datamint.mlflow import set_project from datamint.mlflow.flavors.model import BaseDatamintModel -from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule -from datamint._repr_utils import render_text_block, render_html_card if TYPE_CHECKING: from albumentations import BaseCompose + from datamint.entities import Project _LOGGER = logging.getLogger(__name__) @@ -75,15 +76,15 @@ class BaseTrainer(ABC): def __init__( self, dataset: DatamintBaseDataset | None = None, - project: 'str | Project | None' = None, + project: str | Project | None = None, *, dataset_kwargs: dict[str, Any] | None = None, model: L.LightningModule | type[L.LightningModule] | None = None, loss_fn: nn.Module | None = None, batch_size: int = 16, num_workers: int = 4, - train_transform: 'BaseCompose | None' = None, - eval_transform: 'BaseCompose | None' = None, + train_transform: BaseCompose | None = None, + eval_transform: BaseCompose | None = None, split_as_of_timestamp: str | None = None, max_epochs: int = 1, early_stopping_patience: int | None = 10, @@ -195,12 +196,12 @@ def _create_lightning_trainer( callbacks = self._build_default_callbacks(register_model=register_model) + list(self._build_callbacks()) logger = self._build_logger(run_id=run_id) - trainer_params: dict[str, Any] = dict( - max_epochs=self.max_epochs, - logger=logger, - callbacks=callbacks, - accelerator='auto', - ) + trainer_params: dict[str, Any] = { + "max_epochs": self.max_epochs, + "logger": logger, + "callbacks": callbacks, + "accelerator": 'auto', + } trainer_params.update(self.trainer_kwargs) if override_params is not None: trainer_params.update(override_params) @@ -290,7 +291,7 @@ def test(self, register_model: bool = True) -> list[Mapping[str, float]]: return self._lightning_trainer.test(self.model, datamodule=self.datamodule) - def _upload_test_predictions(self, predict_model: 'BaseDatamintModel | None' = None) -> None: + def _upload_test_predictions(self, predict_model: BaseDatamintModel | None = None) -> None: """Run inference on the test split and upload predictions as annotations.""" if predict_model is None: _LOGGER.debug("No deployable model available; skipping test prediction upload.") @@ -333,7 +334,7 @@ def _upload_test_predictions(self, predict_model: 'BaseDatamintModel | None' = N @abstractmethod def _build_dataset( self, - project: 'str | Project', + project: str | Project, **kwargs: Any ) -> DatamintBaseDataset: """Build the appropriate dataset for this task.""" @@ -357,7 +358,7 @@ def model(self) -> L.LightningModule: try: eval_tf = self._user_eval_transform or self._eval_transform() - setattr(model, 'transform', eval_tf) + model.transform = eval_tf except NotImplementedError: pass @@ -372,12 +373,12 @@ def _build_model( raise NotImplementedError("Subclasses must implement _build_model() when no user model is provided.") @abstractmethod - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: """Return the training augmentation pipeline.""" ... @abstractmethod - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: """Return the eval/test transform pipeline.""" ... @@ -430,8 +431,8 @@ def datamodule(self) -> DatamintDataModule: def _build_datamodule( self, dataset: DatamintBaseDataset, - train_transform: 'BaseCompose', - eval_transform: 'BaseCompose', + train_transform: BaseCompose, + eval_transform: BaseCompose, ) -> DatamintDataModule: return DatamintDataModule( dataset, @@ -444,9 +445,13 @@ def _build_datamodule( ) def _build_default_callbacks(self, *, register_model: bool = True) -> list: - from datamint.mlflow.lightning.callbacks import MLFlowPyTorchModelCheckpoint, MLFlowDatamintModelCheckpoint from mlflow.pyfunc.model import PythonModel + from datamint.mlflow.lightning.callbacks import ( + MLFlowDatamintModelCheckpoint, + MLFlowPyTorchModelCheckpoint, + ) + has_val = self.datamodule.has_val_split if has_val: metric_name, mode = self._monitor_metric() @@ -463,14 +468,14 @@ def _build_default_callbacks(self, *, register_model: bool = True) -> list: _LOGGER.debug("Using %s for model checkpointing with monitor='%s' mode='%s'", checkpoint_cls.__name__, metric_name, mode) - checkpoint_kwargs: dict[str, Any] = dict( - monitor=metric_name, - mode=mode, - save_top_k=1, - model_name=model_name, - register_model_on='test', - log_model_metrics=True, - ) + checkpoint_kwargs: dict[str, Any] = { + "monitor": metric_name, + "mode": mode, + "save_top_k": 1, + "model_name": model_name, + "register_model_on": 'test', + "log_model_metrics": True, + } if checkpoint_cls is MLFlowDatamintModelCheckpoint: checkpoint_kwargs['annotation_specs'] = self._build_annotation_specs() @@ -509,14 +514,14 @@ def _build_logger(self, run_id: str | None = None): dataset = self.datamodule.get_mlflow_dataset_split('test') if dataset is None: dataset = self.datamodule.get_mlflow_dataset() - setattr(mlflow_logger, '_mlflow_dataset', dataset) + mlflow_logger._mlflow_dataset = dataset return mlflow_logger class _LogDatasetSplitsCallback(L.Callback): """Lightning callback to retrieve resolved dataset splits from the datamodule after setup().""" - LIGHTNING_STAGE_TO_DATAMINT_SPLIT = { + LIGHTNING_STAGE_TO_DATAMINT_SPLIT: ClassVar[dict[str, str]] = { 'fit': 'train', 'validate': 'val', 'test': 'test', @@ -526,7 +531,7 @@ def __init__(self, dttrainer: BaseTrainer) -> None: super().__init__() self.dttrainer = dttrainer - def setup(self, trainer: "L.Trainer", pl_module: "L.LightningModule", stage: str) -> None: + def setup(self, trainer: L.Trainer, pl_module: L.LightningModule, stage: str) -> None: split = self.LIGHTNING_STAGE_TO_DATAMINT_SPLIT.get(stage) if split is None: diff --git a/datamint/lightning/trainers/classification_trainer.py b/datamint/lightning/trainers/classification_trainer.py index 17c077bb..54a04472 100644 --- a/datamint/lightning/trainers/classification_trainer.py +++ b/datamint/lightning/trainers/classification_trainer.py @@ -2,21 +2,22 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, TYPE_CHECKING +from functools import partial +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn from datamint.dataset import ImageDataset -from functools import partial - from datamint.entities.annotations.annotation_spec import CategoryAnnotationSpec from datamint.entities.annotations.types import AnnotationType -from .lightning_modules import ClassificationModule + from .base_trainer import BaseTrainer +from .lightning_modules import ClassificationModule if TYPE_CHECKING: from albumentations import BaseCompose + from datamint.entities import Project @@ -107,13 +108,13 @@ def _extra_repr_fields(self) -> list[tuple[str, str]]: # ── Template hooks ────────────────────────────────────────── - def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> ImageDataset: - default_params = dict( - return_segmentations=False, - include_unannotated=False, - image_categories_merge_strategy='mode', - allow_external_annotations=True, - ) + def _build_dataset(self, project: str | Project, **kwargs: Any) -> ImageDataset: + default_params = { + 'return_segmentations': False, + 'include_unannotated': False, + 'image_categories_merge_strategy': 'mode', + 'allow_external_annotations': True, + } dataset_params = {**default_params, **kwargs} return ImageDataset( project=project, @@ -142,7 +143,7 @@ def _build_resize_transform(self): return A.NoOp() return A.Resize(*self.image_size) - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 @@ -155,7 +156,7 @@ def _train_transform(self) -> 'BaseCompose': ToTensorV2(), ]) - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 diff --git a/datamint/lightning/trainers/detection_trainer.py b/datamint/lightning/trainers/detection_trainer.py index 99732e1e..d1d3b0ba 100644 --- a/datamint/lightning/trainers/detection_trainer.py +++ b/datamint/lightning/trainers/detection_trainer.py @@ -2,14 +2,13 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, TYPE_CHECKING - -from torch import nn +from typing import TYPE_CHECKING, Any from datamint.dataset.image_dataset import ImageDataset, detection_collate_fn -from datamint.lightning.datamodule import DatamintDataModule from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.annotations.types import AnnotationType +from datamint.lightning.datamodule import DatamintDataModule + from .base_trainer import BaseTrainer if TYPE_CHECKING: @@ -32,7 +31,7 @@ class DetectionTrainer(BaseTrainer): def _build_dataset( self, - project: 'str | Project', + project: str | Project, **kwargs: Any, ) -> ImageDataset: return ImageDataset(project=project, return_boxes=True, **kwargs) diff --git a/datamint/lightning/trainers/lightning_modules/__init__.py b/datamint/lightning/trainers/lightning_modules/__init__.py index 0796730a..2b02fba8 100644 --- a/datamint/lightning/trainers/lightning_modules/__init__.py +++ b/datamint/lightning/trainers/lightning_modules/__init__.py @@ -1,7 +1,13 @@ from .base import DatamintLightningModule -from .segmentation_module import SegmentationModule -from .segmentation_modules import SMPSegmentationModule, UNetPPModule, DeepLabV3PlusModule, TransUNetModule, UNETRPPModule from .classification_module import ClassificationModule from .detection_modules import YOLOXModule +from .segmentation_module import SegmentationModule +from .segmentation_modules import ( + DeepLabV3PlusModule, + SMPSegmentationModule, + TransUNetModule, + UNetPPModule, + UNETRPPModule, +) -__all__ = ["DatamintLightningModule", "SegmentationModule", "SMPSegmentationModule", "UNetPPModule", "DeepLabV3PlusModule", "TransUNetModule", "UNETRPPModule", "ClassificationModule", "YOLOXModule"] +__all__ = ["ClassificationModule", "DatamintLightningModule", "DeepLabV3PlusModule", "SMPSegmentationModule", "SegmentationModule", "TransUNetModule", "UNETRPPModule", "UNetPPModule", "YOLOXModule"] diff --git a/datamint/lightning/trainers/lightning_modules/base.py b/datamint/lightning/trainers/lightning_modules/base.py index 7cad5784..e7aa125e 100644 --- a/datamint/lightning/trainers/lightning_modules/base.py +++ b/datamint/lightning/trainers/lightning_modules/base.py @@ -2,18 +2,18 @@ from __future__ import annotations import logging -from typing import Any, TYPE_CHECKING +import time +from typing import TYPE_CHECKING, Any +import albumentations as A import lightning as L -from torch import Tensor import torch - -from datamint.mlflow.flavors.model import BaseDatamintModel, ModelSettings -from mlflow.pyfunc.model import PythonModelContext -import time from mlflow.entities import Metric +from mlflow.pyfunc.model import PythonModelContext from mlflow.tracking import MlflowClient -import albumentations as A +from torch import Tensor + +from datamint.mlflow.flavors.model import BaseDatamintModel, ModelSettings if TYPE_CHECKING: from datamint.mlflow.data import DatamintMLflowDataset @@ -144,6 +144,7 @@ def _flush_sample_metrics_to_mlflow(self) -> None: return import mlflow + from datamint.mlflow.models import _get_MLFlowLogger logger = _get_MLFlowLogger(self.trainer) diff --git a/datamint/lightning/trainers/lightning_modules/classification_module.py b/datamint/lightning/trainers/lightning_modules/classification_module.py index 44e983a9..082d86fb 100644 --- a/datamint/lightning/trainers/lightning_modules/classification_module.py +++ b/datamint/lightning/trainers/lightning_modules/classification_module.py @@ -1,19 +1,21 @@ """LightningModule wrapper for image classification tasks.""" from __future__ import annotations -from collections.abc import Callable import inspect -from typing import Any import warnings +from collections.abc import Callable +from typing import Any import albumentations as A +import numpy as np import torch +from albumentations.pytorch import ToTensorV2 from torch import Tensor, nn from torchmetrics import MetricCollection -import numpy as np -from albumentations.pytorch import ToTensorV2 + from datamint.entities.annotations import ImageClassification from datamint.mlflow.flavors.task_type import TaskType + from .base import DatamintLightningModule diff --git a/datamint/lightning/trainers/lightning_modules/detection_modules/yolox_module.py b/datamint/lightning/trainers/lightning_modules/detection_modules/yolox_module.py index b9d8c978..9e94321a 100644 --- a/datamint/lightning/trainers/lightning_modules/detection_modules/yolox_module.py +++ b/datamint/lightning/trainers/lightning_modules/detection_modules/yolox_module.py @@ -8,6 +8,7 @@ from torch import Tensor from datamint.mlflow.flavors.task_type import TaskType + from ..base import DatamintLightningModule _LOGGER = logging.getLogger(__name__) @@ -212,6 +213,7 @@ def predict_image(self, model_input: Any, **kwargs: Any) -> list: import numpy as np from albumentations.pytorch import ToTensorV2 from yolox.utils import postprocess as yolox_postprocess + from datamint.entities.annotations import BoxAnnotation self.eval() diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_module.py b/datamint/lightning/trainers/lightning_modules/segmentation_module.py index ef6d4d5c..e9d08260 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_module.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_module.py @@ -1,21 +1,21 @@ """LightningModule wrapper for segmentation tasks.""" from __future__ import annotations +import inspect +import warnings from abc import abstractmethod from collections.abc import Callable from typing import Any -import inspect -import warnings +import albumentations as A import torch +from albumentations.pytorch import ToTensorV2 from torch import Tensor, nn from torchmetrics import MetricCollection from datamint.mlflow.flavors.task_type import TaskType + from .base import DatamintLightningModule -from datamint.mlflow.flavors.prediction_router import prediction_mode -import albumentations as A -from albumentations.pytorch import ToTensorV2 class SegmentationModule(DatamintLightningModule): @@ -39,12 +39,14 @@ class SegmentationModule(DatamintLightningModule): def __init__( self, loss_fn: nn.Module | None = None, - metrics_factories: dict[str, Callable[[], Any]] = {}, + metrics_factories: dict[str, Callable[[], Any]] | None = None, class_names: list[str] | None = None, # image_size: tuple[int, int], transform: A.BasicTransform | A.BaseCompose | None = None, lr: float = 1e-4, ) -> None: + if metrics_factories is None: + metrics_factories = {} super().__init__(transform=transform) self.save_hyperparameters(ignore=['loss_fn', 'metrics_factories', 'transform']) self.class_names = class_names @@ -193,6 +195,7 @@ def predict_image(self, model_input, compute_uncertainty: bool = False, **kwargs """ import cv2 import numpy as np + from datamint.entities.annotations import ImageSegmentation from datamint.utils.uncertainty import segmentation_uncertainty diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/__init__.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/__init__.py index f1512342..daf4f96a 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/__init__.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/__init__.py @@ -1,7 +1,7 @@ -from .smp_module import SMPSegmentationModule -from .unetpp import UNetPPModule from .deeplabv3plus import DeepLabV3PlusModule +from .smp_module import SMPSegmentationModule from .transunet import TransUNetModule +from .unetpp import UNetPPModule from .unetrpp import UNETRPPModule -__all__ = ["SMPSegmentationModule", "UNetPPModule", "DeepLabV3PlusModule", "TransUNetModule", "UNETRPPModule"] +__all__ = ["DeepLabV3PlusModule", "SMPSegmentationModule", "TransUNetModule", "UNETRPPModule", "UNetPPModule"] diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/deeplabv3plus.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/deeplabv3plus.py index 9f54aca7..8c629072 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/deeplabv3plus.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/deeplabv3plus.py @@ -1,14 +1,14 @@ """DeepLab v3+ segmentation module.""" from __future__ import annotations + from collections.abc import Callable from typing import Any import albumentations as A +import segmentation_models_pytorch as smp from torch import Tensor, nn from typing_extensions import override -import segmentation_models_pytorch as smp - from .smp_module import SMPSegmentationModule @@ -36,7 +36,7 @@ def __init__( in_channels: int, num_classes: int, loss_fn: nn.Module | None = None, - metrics_factories: dict[str, Callable[[], Any]] = {}, + metrics_factories: dict[str, Callable[[], Any]] | None = None, class_names: list[str] | None = None, image_size: tuple[int, int] | None = None, lr: float = 1e-4, @@ -45,6 +45,8 @@ def __init__( decoder_atrous_rates: tuple[int, int, int] = (12, 24, 36), transform: A.BasicTransform | A.BaseCompose | None = None, ) -> None: + if metrics_factories is None: + metrics_factories = {} self.in_channels = in_channels self.num_classes = num_classes self.encoder_name = encoder_name diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/smp_module.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/smp_module.py index d13c073d..4ee2e286 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/smp_module.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/smp_module.py @@ -1,8 +1,6 @@ """Base segmentation module for ``segmentation_models_pytorch`` architectures.""" from __future__ import annotations -from typing import Any - from ..segmentation_module import SegmentationModule @@ -14,4 +12,3 @@ class SMPSegmentationModule(SegmentationModule): concrete SMP model. """ - pass diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/transunet.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/transunet.py index 618b1fde..68cd68c7 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/transunet.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/transunet.py @@ -11,9 +11,8 @@ import albumentations as A import torch -import torch.nn as nn import torch.nn.functional as F -from torch import Tensor +from torch import Tensor, nn from typing_extensions import override from ..segmentation_module import SegmentationModule @@ -79,7 +78,7 @@ def __init__( pretrained: bool = True, ) -> None: super().__init__() - import timm + import timm cfg = _VARIANT_CONFIGS[variant] hidden_dim: int = cfg['hidden_dim'] @@ -105,7 +104,7 @@ def __init__( self._norm = vit.norm # final LayerNorm # CUP decoder: 4 blocks, first 2 fuse CNN skip features. - decoder_in_ch = [hidden_dim] + dec_ch[:-1] + decoder_in_ch = [hidden_dim, *dec_ch[:-1]] self._decoder = nn.ModuleList([ _DecoderBlock(in_ch, out_ch, s_ch) for in_ch, out_ch, s_ch in zip(decoder_in_ch, dec_ch, skip_ch) @@ -183,7 +182,7 @@ def __init__( in_channels: int, num_classes: int, loss_fn: nn.Module | None = None, - metrics_factories: dict[str, Callable[[], Any]] = {}, + metrics_factories: dict[str, Callable[[], Any]] | None = None, class_names: list[str] | None = None, image_size: tuple[int, int] | None = None, lr: float = 1e-4, @@ -191,6 +190,8 @@ def __init__( pretrained: bool = True, transform: A.BasicTransform | A.BaseCompose | None = None, ) -> None: + if metrics_factories is None: + metrics_factories = {} if variant not in _VARIANT_CONFIGS: raise ValueError(f"Unknown variant {variant!r}. Choose from {list(_VARIANT_CONFIGS)}") diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetpp.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetpp.py index 53b8860f..3d471871 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetpp.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetpp.py @@ -1,13 +1,13 @@ """UNet++ segmentation module.""" from __future__ import annotations -from typing import Any -from typing_extensions import override -from collections.abc import Callable -from torch import Tensor, nn +from collections.abc import Callable +from typing import Any import albumentations as A import segmentation_models_pytorch as smp +from torch import Tensor, nn +from typing_extensions import override from .smp_module import SMPSegmentationModule @@ -25,7 +25,7 @@ def __init__( in_channels: int, num_classes: int, loss_fn: nn.Module | None = None, - metrics_factories: dict[str, Callable[[], Any]] = {}, + metrics_factories: dict[str, Callable[[], Any]] | None = None, class_names: list[str] | None = None, image_size: tuple[int, int] | None = None, lr: float = 1e-4, @@ -33,6 +33,8 @@ def __init__( encoder_weights: str | None = 'imagenet', transform: A.BasicTransform | A.BaseCompose | None = None, ) -> None: + if metrics_factories is None: + metrics_factories = {} self.in_channels = in_channels self.num_classes = num_classes self.encoder_name = encoder_name diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py index 38dd224a..9ad6d4c1 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py @@ -13,17 +13,15 @@ from collections.abc import Callable from typing import Any +import albumentations as A +import numpy as np import torch -import torch.nn as nn import torch.nn.functional as F -from torch import Tensor +from torch import Tensor, nn from typing_extensions import override -import albumentations as A - from ..segmentation_module import SegmentationModule - # --------------------------------------------------------------------------- # Internal building blocks — not exported # --------------------------------------------------------------------------- @@ -394,7 +392,7 @@ def _proj_feat(self, x: Tensor) -> Tensor: Source: https://github.com/Amshaker/unetr_plus_plus/blob/main/unetr_pp/network_architecture/synapse/unetr_pp_synapse.py """ - B, N, C = x.shape + B, _N, C = x.shape D, H, W = self.feat_size return x.reshape(B, D, H, W, C).permute(0, 4, 1, 2, 3).contiguous() @@ -452,7 +450,7 @@ def __init__( num_classes: int, img_size: tuple[int, int, int], loss_fn: nn.Module | None = None, - metrics_factories: dict[str, Callable[[], Any]] = {}, + metrics_factories: dict[str, Callable[[], Any]] | None = None, class_names: list[str] | None = None, lr: float = 1e-4, feature_size: int = 16, @@ -463,6 +461,8 @@ def __init__( transform: A.BasicTransform | A.BaseCompose | None = None, ) -> None: + if metrics_factories is None: + metrics_factories = {} self.in_channels = in_channels self.num_classes = num_classes self.feature_size = feature_size @@ -522,7 +522,7 @@ def _random_crop_3d(self, image: Tensor, mask: Tensor) -> tuple[Tensor, Tensor]: ) @staticmethod - def _to_float_tensor(x: 'Tensor | np.ndarray') -> Tensor: + def _to_float_tensor(x: Tensor | np.ndarray) -> Tensor: """Convert numpy arrays to float tensors. ToTensorV2 in albumentations only converts the 'image' key, not the @@ -586,7 +586,7 @@ def _sliding_window_inference(self, volume: Tensor) -> Tensor: """ device = next(self.parameters()).device volume = volume.float().to(device) - B, C, D0, H0, W0 = volume.shape + B, _C, D0, H0, W0 = volume.shape pd, ph, pw = self.patch_crop_size pad_d, pad_h, pad_w = max(0, pd - D0), max(0, ph - H0), max(0, pw - W0) @@ -661,8 +661,9 @@ def predict_volume( import numpy as np from albumentations.pytorch import ToTensorV2 from medimgkit.readers import read_array_normalized + from datamint.entities.annotations import VolumeSegmentation - from datamint.utils.uncertainty import segmentation_uncertainty, pool_top_k + from datamint.utils.uncertainty import pool_top_k, segmentation_uncertainty transform = self.transform or A.Compose([A.Normalize(), ToTensorV2()]) device = self.inference_device diff --git a/datamint/lightning/trainers/seg2d_trainer.py b/datamint/lightning/trainers/seg2d_trainer.py index 8232bb85..d369a04e 100644 --- a/datamint/lightning/trainers/seg2d_trainer.py +++ b/datamint/lightning/trainers/seg2d_trainer.py @@ -2,12 +2,12 @@ from __future__ import annotations from collections.abc import Mapping, Sequence -from typing import Any, Literal, TYPE_CHECKING, cast -from typing_extensions import override +from typing import TYPE_CHECKING, Any, Literal, cast -import lightning as L import albumentations as A +import lightning as L from albumentations.pytorch import ToTensorV2 +from typing_extensions import override from datamint.dataset import ImageDataset, SlicedVolumeDataset from datamint.utils.nifti_utils import metadata_to_nifti_obj @@ -16,9 +16,9 @@ if TYPE_CHECKING: from albumentations import BaseCompose + from medimgkit import ViewPlane from nibabel.spatialimages import SpatialImage from pydicom import Dataset as DicomDataset - from medimgkit import ViewPlane from datamint.entities import Project, Resource @@ -61,7 +61,7 @@ def __init__( self, *, image_size: int | tuple[int, int] | None = None, - slice_axis: 'ViewPlane | int | None' = None, + slice_axis: ViewPlane | int | None = None, model: L.LightningModule | type[L.LightningModule] | None = None, in_channels: int = 3, trainer_kwargs: dict[str, Any] | None = None, @@ -71,7 +71,7 @@ def __init__( trainer_kwargs=trainer_kwargs, **kwargs) self.in_channels = in_channels - self.slice_axis: 'ViewPlane | int | None' = slice_axis + self.slice_axis: ViewPlane | int | None = slice_axis if isinstance(image_size, int): self.image_size = (image_size, image_size) else: @@ -82,13 +82,13 @@ def _extra_repr_fields(self) -> list[tuple[str, str]]: image_size = f"{self.image_size[0]}×{self.image_size[1]}" if self.image_size else "auto (no resize)" return [*super()._extra_repr_fields(), ("Image size", image_size)] - def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> ImageDataset | SlicedVolumeDataset: - default_params = dict( - return_as_semantic_segmentation=True, - semantic_seg_merge_strategy='union', - allow_external_annotations=True, - include_unannotated=False, - ) + def _build_dataset(self, project: str | Project, **kwargs: Any) -> ImageDataset | SlicedVolumeDataset: + default_params = { + 'return_as_semantic_segmentation': True, + 'semantic_seg_merge_strategy': 'union', + 'allow_external_annotations': True, + 'include_unannotated': False, + } dataset_params = {**default_params, **kwargs} dataset = ImageDataset( project=project, @@ -120,7 +120,7 @@ def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> ImageDatase _LOGGER.info("Project contains 2D images; using ImageDataset.") return dataset - def _classify_resource(self, resource: 'Resource') -> str: + def _classify_resource(self, resource: Resource) -> str: if resource.is_video(): return 'video' @@ -137,13 +137,13 @@ def _classify_resource(self, resource: 'Resource') -> str: return getattr(resource, 'kind', 'unknown') @staticmethod - def _get_resource_depth(resource: 'Resource') -> int | None: + def _get_resource_depth(resource: Resource) -> int | None: try: return resource.get_depth() except Exception: return None - def _infer_slice_axis(self, resources: Sequence['Resource']) -> SliceAxisName: + def _infer_slice_axis(self, resources: Sequence[Resource]) -> SliceAxisName: for resource in resources: if self._classify_resource(resource) != 'volume': continue @@ -154,7 +154,7 @@ def _infer_slice_axis(self, resources: Sequence['Resource']) -> SliceAxisName: return 'axial' - def _infer_slice_axis_from_resource(self, resource: 'Resource') -> SliceAxisName | None: + def _infer_slice_axis_from_resource(self, resource: Resource) -> SliceAxisName | None: if resource.is_nifti(): nifti_image = self._nifti_image_from_metadata(resource) if nifti_image is not None: @@ -176,7 +176,7 @@ def _infer_slice_axis_from_resource(self, resource: 'Resource') -> SliceAxisName return None @staticmethod - def _nifti_image_from_metadata(resource: 'Resource') -> 'SpatialImage | None': + def _nifti_image_from_metadata(resource: Resource) -> SpatialImage | None: metadata = getattr(resource, 'metadata', None) if not isinstance(metadata, dict): return None @@ -186,7 +186,7 @@ def _nifti_image_from_metadata(resource: 'Resource') -> 'SpatialImage | None': except Exception: return None - def _infer_slice_axis_from_nifti(self, nifti_image: 'SpatialImage') -> SliceAxisName | None: + def _infer_slice_axis_from_nifti(self, nifti_image: SpatialImage) -> SliceAxisName | None: from medimgkit import nifti_utils plane_sizes = { @@ -201,7 +201,7 @@ def _infer_slice_axis_from_nifti(self, nifti_image: 'SpatialImage') -> SliceAxis } return self._choose_slice_axis(plane_sizes, plane_spacings) - def _infer_slice_axis_from_dicom(self, dataset: 'DicomDataset') -> SliceAxisName | None: + def _infer_slice_axis_from_dicom(self, dataset: DicomDataset) -> SliceAxisName | None: from medimgkit import dicom_utils pixel_spacing = self._coerce_spacing_pair(getattr(dataset, 'PixelSpacing', None)) @@ -280,7 +280,7 @@ def _build_resize_transform(self): return A.Resize(*self.image_size) @override - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: return A.Compose([ self._build_resize_transform(), A.ToRGB(), @@ -292,7 +292,7 @@ def _train_transform(self) -> 'BaseCompose': ]) @override - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: return A.Compose([ self._build_resize_transform(), A.ToRGB(), diff --git a/datamint/lightning/trainers/seg3d_trainer.py b/datamint/lightning/trainers/seg3d_trainer.py index 52530ddb..6e9eaabc 100644 --- a/datamint/lightning/trainers/seg3d_trainer.py +++ b/datamint/lightning/trainers/seg3d_trainer.py @@ -1,10 +1,7 @@ """3-D semantic segmentation trainer (slice-based).""" from __future__ import annotations -from typing import Any, TYPE_CHECKING - -import lightning as L -from torch import nn +from typing import TYPE_CHECKING, Any from datamint.dataset import VolumeDataset @@ -12,6 +9,7 @@ if TYPE_CHECKING: from albumentations import BaseCompose + from datamint.entities import Project class SemanticSegmentation3DTrainer(SegmentationTrainer): @@ -72,13 +70,13 @@ def _extra_repr_fields(self) -> list[tuple[str, str]]: # ── Template hooks ────────────────────────────────────────── - def _build_dataset(self, project: 'str | Project', **kwargs: Any): - default_params = dict( - return_as_semantic_segmentation=True, - semantic_seg_merge_strategy='union', - allow_external_annotations=True, - include_unannotated=False, - ) + def _build_dataset(self, project: str | Project, **kwargs: Any): + default_params = { + 'return_as_semantic_segmentation': True, + 'semantic_seg_merge_strategy': 'union', + 'allow_external_annotations': True, + 'include_unannotated': False, + } dataset_params = {**default_params, **kwargs} vol_ds = VolumeDataset( @@ -95,7 +93,7 @@ def _build_resize_transform(self): return A.NoOp() return A.Resize(*self.image_size) - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 @@ -107,7 +105,7 @@ def _train_transform(self) -> 'BaseCompose': ToTensorV2(), ]) - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 diff --git a/datamint/lightning/trainers/segmentation_trainer.py b/datamint/lightning/trainers/segmentation_trainer.py index 088e5ef1..0e492662 100644 --- a/datamint/lightning/trainers/segmentation_trainer.py +++ b/datamint/lightning/trainers/segmentation_trainer.py @@ -2,14 +2,15 @@ from __future__ import annotations from collections.abc import Callable - from functools import partial + import torch import torch.nn.functional as F from torch import nn from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.annotations.types import AnnotationType + from .base_trainer import BaseTrainer diff --git a/datamint/lightning/trainers/specialized/__init__.py b/datamint/lightning/trainers/specialized/__init__.py index 6277cb57..09b74b16 100644 --- a/datamint/lightning/trainers/specialized/__init__.py +++ b/datamint/lightning/trainers/specialized/__init__.py @@ -1,18 +1,17 @@ -from .unetpp import UNetPPTrainer from .deeplabv3plus import DeepLabV3PlusTrainer +from .efficientnetv2 import EfficientNetV2Trainer +from .nnunet.trainer import NNUNetTrainer from .transunet import TransUNetTrainer +from .unetpp import UNetPPTrainer from .unetrpp import UNETRPPTrainer -from .nnunet.trainer import NNUNetTrainer from .yolox import YOLOXTrainer -from .efficientnetv2 import EfficientNetV2Trainer - __all__ = [ - "UNetPPTrainer", "DeepLabV3PlusTrainer", + "EfficientNetV2Trainer", + "NNUNetTrainer", "TransUNetTrainer", "UNETRPPTrainer", - "NNUNetTrainer", + "UNetPPTrainer", "YOLOXTrainer", - "EfficientNetV2Trainer", ] diff --git a/datamint/lightning/trainers/specialized/deeplabv3plus.py b/datamint/lightning/trainers/specialized/deeplabv3plus.py index 00746c33..4a06b6d9 100644 --- a/datamint/lightning/trainers/specialized/deeplabv3plus.py +++ b/datamint/lightning/trainers/specialized/deeplabv3plus.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn @@ -10,10 +10,13 @@ if TYPE_CHECKING: from albumentations import BaseCompose + from medimgkit import ViewPlane + from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project - from medimgkit import ViewPlane - from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule + from datamint.lightning.trainers.lightning_modules.base import ( + DatamintLightningModule, + ) class DeepLabV3PlusTrainer(SemanticSegmentation2DTrainer): diff --git a/datamint/lightning/trainers/specialized/efficientnetv2.py b/datamint/lightning/trainers/specialized/efficientnetv2.py index afb2e3ce..c5b6e393 100644 --- a/datamint/lightning/trainers/specialized/efficientnetv2.py +++ b/datamint/lightning/trainers/specialized/efficientnetv2.py @@ -1,7 +1,8 @@ """EfficientNetV2 image classification trainer.""" from __future__ import annotations -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any + from typing_extensions import override from ..classification_trainer import ImageClassificationTrainer @@ -39,7 +40,7 @@ def __init__( super().__init__(architecture=architecture, image_size=image_size, **kwargs) @override - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 diff --git a/datamint/lightning/trainers/specialized/nnunet/_nnunet_trainer_bridge.py b/datamint/lightning/trainers/specialized/nnunet/_nnunet_trainer_bridge.py index de10d408..5b1a2f14 100644 --- a/datamint/lightning/trainers/specialized/nnunet/_nnunet_trainer_bridge.py +++ b/datamint/lightning/trainers/specialized/nnunet/_nnunet_trainer_bridge.py @@ -1,11 +1,11 @@ from __future__ import annotations +import importlib.metadata as _importlib_metadata import json import logging -import mlflow from pathlib import Path -import importlib.metadata as _importlib_metadata +import mlflow _LOGGER = logging.getLogger(__name__) @@ -31,8 +31,7 @@ def _parse_version(v: str) -> tuple[int, ...]: ) # ────────────────────────────────────────────────────────────────────────────── -from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer - +from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer # Keys logged by nnUNet that are not useful as MLflow metrics. _SKIP_METRIC_KEYS = frozenset({'epoch_start_timestamps', 'epoch_end_timestamps'}) diff --git a/datamint/lightning/trainers/specialized/nnunet/data_export.py b/datamint/lightning/trainers/specialized/nnunet/data_export.py index 89497fdf..4a9f43d6 100644 --- a/datamint/lightning/trainers/specialized/nnunet/data_export.py +++ b/datamint/lightning/trainers/specialized/nnunet/data_export.py @@ -4,6 +4,7 @@ import logging import warnings from pathlib import Path + import nibabel as nib import numpy as np @@ -69,7 +70,7 @@ def _export_image(self, resource, case_id: str, split: str) -> Path: def _merge_segmentations( self, segs, - name_to_idx: 'dict[str, int] | None' = None, + name_to_idx: dict[str, int] | None = None, ) -> np.ndarray: """Merge N segmentation masks into one int32 label map. diff --git a/datamint/lightning/trainers/specialized/nnunet/data_import.py b/datamint/lightning/trainers/specialized/nnunet/data_import.py index 0e275533..f5d92a8f 100644 --- a/datamint/lightning/trainers/specialized/nnunet/data_import.py +++ b/datamint/lightning/trainers/specialized/nnunet/data_import.py @@ -3,6 +3,7 @@ import json import logging from pathlib import Path + import nibabel as nib import numpy as np diff --git a/datamint/lightning/trainers/specialized/nnunet/inference_model.py b/datamint/lightning/trainers/specialized/nnunet/inference_model.py index 026bfe2e..28c13785 100644 --- a/datamint/lightning/trainers/specialized/nnunet/inference_model.py +++ b/datamint/lightning/trainers/specialized/nnunet/inference_model.py @@ -66,8 +66,8 @@ def load_context(self, context) -> None: """ super().load_context(context) os.environ.setdefault('nnUNet_extTrainer', str(Path(__file__).parent)) - import torch import nnunetv2.inference.predict_from_raw_data as _pred_mod + import torch bundle_path = context.artifacts['nnunet_bundle'] predictor = _pred_mod.nnUNetPredictor(device=torch.device(self.inference_device)) diff --git a/datamint/lightning/trainers/specialized/nnunet/trainer.py b/datamint/lightning/trainers/specialized/nnunet/trainer.py index 06e3613d..8c7ab19d 100644 --- a/datamint/lightning/trainers/specialized/nnunet/trainer.py +++ b/datamint/lightning/trainers/specialized/nnunet/trainer.py @@ -3,18 +3,20 @@ import logging import os from pathlib import Path -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any import filelock import mlflow import yaml from rich import print as rprint -from datamint.lightning.trainers.base_trainer import BaseTrainer from datamint.dataset.volume_dataset import VolumeDataset -from datamint.lightning.trainers.specialized.nnunet.data_export import DatamintToNNUNetExporter from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.annotations.types import AnnotationType +from datamint.lightning.trainers.base_trainer import BaseTrainer +from datamint.lightning.trainers.specialized.nnunet.data_export import ( + DatamintToNNUNetExporter, +) if TYPE_CHECKING: from datamint.entities import Project @@ -54,7 +56,7 @@ class NNUNetTrainer(BaseTrainer): def __init__( self, dataset=None, - project: 'str | Project | None' = None, + project: str | Project | None = None, *, configuration: str = '3d_fullres', fold: int | str = 0, @@ -100,7 +102,7 @@ def _repr_fields(self) -> list[tuple[str, str]]: # ── BaseTrainer abstract methods bypassed by nnUNet ─────────────────────── - def _build_dataset(self, project: 'str | Project', **kwargs) -> VolumeDataset: + def _build_dataset(self, project: str | Project, **kwargs) -> VolumeDataset: return VolumeDataset(project=project, **kwargs) def _build_annotation_specs(self) -> list[AnnotationSpec]: @@ -239,18 +241,20 @@ def _run_fingerprint_and_plan(self, dataset_id: int) -> None: plans_file = preprocessed_dataset_dir / 'nnUNetPlans.json' if fp_file.exists() and plans_file.exists(): - rprint(f"[green]✓[/green] Fingerprinting and planning already done for dataset — skipping.") + rprint("[green]✓[/green] Fingerprinting and planning already done for dataset — skipping.") return from nnunetv2.experiment_planning.dataset_fingerprint.fingerprint_extractor import ( DatasetFingerprintExtractor, ) - from nnunetv2.experiment_planning.experiment_planners.default_experiment_planner import ExperimentPlanner + from nnunetv2.experiment_planning.experiment_planners.default_experiment_planner import ( + ExperimentPlanner, + ) - rprint(f"[bold]→[/bold] Running dataset fingerprinting for dataset…") + rprint("[bold]→[/bold] Running dataset fingerprinting for dataset…") DatasetFingerprintExtractor(dataset_id, num_processes=8).run() - rprint(f"[bold]→[/bold] Running experiment planning for dataset…") + rprint("[bold]→[/bold] Running experiment planning for dataset…") ExperimentPlanner(dataset_id, gpu_memory_target_in_gb=8.0).plan_experiment() if not fp_file.exists(): @@ -324,6 +328,7 @@ def _build_nnunet_trainer(self, dataset_id: int): call ``run_training()``. """ import json as _json + from datamint.lightning.trainers.specialized.nnunet._nnunet_trainer_bridge import ( _DatamintNNUNetTrainer, ) @@ -389,7 +394,7 @@ def _run_prediction(self, dataset_id: int, bridge) -> Path | None: rprint(f"[green]✓[/green] Predictions written to {pred_dir}") return pred_dir - def _import_predictions(self, dataset_id: int, pred_dir: 'Path | None') -> None: + def _import_predictions(self, dataset_id: int, pred_dir: Path | None) -> None: """Upload nnUNet test predictions to Datamint as volume annotations. Reads the per-class label map from ``dataset.json`` and delegates @@ -405,6 +410,7 @@ def _import_predictions(self, dataset_id: int, pred_dir: 'Path | None') -> None: return import json as _json + from datamint.lightning.trainers.specialized.nnunet.data_import import ( NNUNetToDatamintImporter, ) @@ -516,6 +522,7 @@ def _build_deploy_adapter(self, dataset_id: int, bridge) -> None: """ import json as _json import shutil + import datamint.mlflow.flavors.datamint_flavor as _datamint_flavor from datamint.lightning.trainers.specialized.nnunet.inference_model import ( NNUNetInferenceModel, diff --git a/datamint/lightning/trainers/specialized/transunet.py b/datamint/lightning/trainers/specialized/transunet.py index a7d956b4..6fa4fb7e 100644 --- a/datamint/lightning/trainers/specialized/transunet.py +++ b/datamint/lightning/trainers/specialized/transunet.py @@ -1,19 +1,22 @@ from collections.abc import Callable -from typing import Any, TYPE_CHECKING -from typing_extensions import override +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn +from typing_extensions import override from ..lightning_modules import TransUNetModule from ..seg2d_trainer import SemanticSegmentation2DTrainer if TYPE_CHECKING: from albumentations import BaseCompose + from medimgkit import ViewPlane + from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project - from medimgkit import ViewPlane - from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule + from datamint.lightning.trainers.lightning_modules.base import ( + DatamintLightningModule, + ) class TransUNetTrainer(SemanticSegmentation2DTrainer): diff --git a/datamint/lightning/trainers/specialized/unetpp.py b/datamint/lightning/trainers/specialized/unetpp.py index 72df57d0..539aa459 100644 --- a/datamint/lightning/trainers/specialized/unetpp.py +++ b/datamint/lightning/trainers/specialized/unetpp.py @@ -1,19 +1,22 @@ from collections.abc import Callable -from typing import Any, TYPE_CHECKING -from typing_extensions import override +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn +from typing_extensions import override + from ..lightning_modules import UNetPPModule from ..seg2d_trainer import SemanticSegmentation2DTrainer - if TYPE_CHECKING: from albumentations import BaseCompose + from medimgkit import ViewPlane + from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project - from medimgkit import ViewPlane - from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule + from datamint.lightning.trainers.lightning_modules.base import ( + DatamintLightningModule, + ) class UNetPPTrainer(SemanticSegmentation2DTrainer): diff --git a/datamint/lightning/trainers/specialized/unetrpp.py b/datamint/lightning/trainers/specialized/unetrpp.py index 8f41ea84..c5c1f221 100644 --- a/datamint/lightning/trainers/specialized/unetrpp.py +++ b/datamint/lightning/trainers/specialized/unetrpp.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn @@ -14,7 +14,9 @@ if TYPE_CHECKING: from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project - from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule + from datamint.lightning.trainers.lightning_modules.base import ( + DatamintLightningModule, + ) class UNETRPPTrainer(VolumeSegmentationTrainer): @@ -62,8 +64,8 @@ class UNETRPPTrainer(VolumeSegmentationTrainer): def __init__( self, - dataset: 'DatamintBaseDataset | None' = None, - project: 'str | Project | None' = None, + dataset: DatamintBaseDataset | None = None, + project: str | Project | None = None, *, patch_crop_size: tuple[int, int, int] = (128, 128, 128), feature_size: int = 16, @@ -124,7 +126,7 @@ def _build_model( self, loss_fn: nn.Module, metrics: dict[str, Callable], - ) -> 'DatamintLightningModule': + ) -> DatamintLightningModule: num_classes = len(self.dataset.seglabel_list) if num_classes == 0: raise ValueError( diff --git a/datamint/lightning/trainers/specialized/yolox.py b/datamint/lightning/trainers/specialized/yolox.py index 39ed515a..ff7d5ba4 100644 --- a/datamint/lightning/trainers/specialized/yolox.py +++ b/datamint/lightning/trainers/specialized/yolox.py @@ -1,7 +1,7 @@ """YOLOXTrainer — one-liner detection trainer backed by YOLOX.""" from __future__ import annotations -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any import lightning as L from torch import nn @@ -11,9 +11,12 @@ if TYPE_CHECKING: from albumentations import BaseCompose + from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project - from datamint.lightning.trainers.lightning_modules.base import DatamintLightningModule + from datamint.lightning.trainers.lightning_modules.base import ( + DatamintLightningModule, + ) class YOLOXTrainer(DetectionTrainer): @@ -58,8 +61,8 @@ class YOLOXTrainer(DetectionTrainer): def __init__( self, - dataset: 'DatamintBaseDataset | None' = None, - project: 'str | Project | None' = None, + dataset: DatamintBaseDataset | None = None, + project: str | Project | None = None, *, model_size: str = 's', conf_thre: float = 0.25, @@ -69,8 +72,8 @@ def __init__( loss_fn: nn.Module | None = None, batch_size: int = 8, num_workers: int = 4, - train_transform: 'BaseCompose | None' = None, - eval_transform: 'BaseCompose | None' = None, + train_transform: BaseCompose | None = None, + eval_transform: BaseCompose | None = None, split_as_of_timestamp: str | None = None, max_epochs: int = 50, early_stopping_patience: int | None = 10, @@ -124,7 +127,7 @@ def _extra_repr_fields(self) -> list[tuple[str, str]]: # Transforms # ------------------------------------------------------------------ - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 @@ -143,7 +146,7 @@ def _train_transform(self) -> 'BaseCompose': ToTensorV2(), ], bbox_params=bbox_params) - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: import albumentations as A from albumentations.pytorch import ToTensorV2 @@ -168,7 +171,7 @@ def _build_model( self, loss_fn: nn.Module | None, metrics: dict, - ) -> 'DatamintLightningModule': + ) -> DatamintLightningModule: num_classes = len(self.dataset.box_class_map) if num_classes == 0: raise ValueError( diff --git a/datamint/lightning/trainers/vol_seg_trainer.py b/datamint/lightning/trainers/vol_seg_trainer.py index 742edfec..067d9a49 100644 --- a/datamint/lightning/trainers/vol_seg_trainer.py +++ b/datamint/lightning/trainers/vol_seg_trainer.py @@ -3,24 +3,24 @@ from collections.abc import Callable from functools import partial -from typing import Any, TYPE_CHECKING +from typing import TYPE_CHECKING, Any +import albumentations as A import torch import torch.nn.functional as F -from torch import nn - -import albumentations as A from albumentations.pytorch import ToTensorV2 +from torch import nn from datamint.dataset import VolumeDataset -from datamint.lightning.datamodule import DatamintDataModule - from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.annotations.types import AnnotationType +from datamint.lightning.datamodule import DatamintDataModule + from .segmentation_trainer import SegmentationTrainer if TYPE_CHECKING: from albumentations import BaseCompose + from datamint.dataset.base import DatamintBaseDataset from datamint.entities import Project @@ -96,13 +96,13 @@ def _extra_repr_fields(self) -> list[tuple[str, str]]: # ── Template hooks ─────────────────────────────────────────── - def _build_dataset(self, project: 'str | Project', **kwargs: Any) -> VolumeDataset: - default_params: dict[str, Any] = dict( - return_as_semantic_segmentation=True, - semantic_seg_merge_strategy='union', - allow_external_annotations=True, - include_unannotated=False, - ) + def _build_dataset(self, project: str | Project, **kwargs: Any) -> VolumeDataset: + default_params: dict[str, Any] = { + 'return_as_semantic_segmentation': True, + 'semantic_seg_merge_strategy': 'union', + 'allow_external_annotations': True, + 'include_unannotated': False, + } return VolumeDataset(project=project, **{**default_params, **kwargs}) def _loss(self) -> nn.Module: @@ -117,7 +117,7 @@ def _metrics(self) -> dict[str, Callable]: 'dice': partial(GeneralizedDiceScore, num_classes=num_classes, input_format='one-hot'), } - def _train_transform(self) -> 'BaseCompose': + def _train_transform(self) -> BaseCompose: """Intensity-only transforms (no crop). Spatial crop is done inside the Lightning module's ``training_step`` @@ -129,7 +129,7 @@ def _train_transform(self) -> 'BaseCompose': ToTensorV2(), ]) - def _eval_transform(self) -> 'BaseCompose': + def _eval_transform(self) -> BaseCompose: """Normalise only — no crop. Full volumes are passed to the Lightning module's sliding-window @@ -142,9 +142,9 @@ def _eval_transform(self) -> 'BaseCompose': def _build_datamodule( self, - dataset: 'DatamintBaseDataset', - train_transform: 'BaseCompose', - eval_transform: 'BaseCompose', + dataset: DatamintBaseDataset, + train_transform: BaseCompose, + eval_transform: BaseCompose, ) -> DatamintDataModule: """Use ``eval_batch_size=1`` for val/test. diff --git a/datamint/mlflow/__init__.py b/datamint/mlflow/__init__.py index 7f15a6c3..3989739e 100644 --- a/datamint/mlflow/__init__.py +++ b/datamint/mlflow/__init__.py @@ -1,10 +1,12 @@ # Monkey patch mlflow.tracking._tracking_service.utils.get_tracking_uri -import mlflow.tracking._tracking_service.utils as mlflow_utils -from functools import wraps import logging -from .env_utils import setup_mlflow_environment, ensure_mlflow_configured +from functools import wraps from typing import TYPE_CHECKING +import mlflow.tracking._tracking_service.utils as mlflow_utils + +from .env_utils import ensure_mlflow_configured, setup_mlflow_environment + _LOGGER = logging.getLogger(__name__) # Store reference to original function @@ -48,12 +50,12 @@ def _configure_mlflow_loggers(): if _ALREADY_CONFIGURED_LOGGING: return - from mlflow.environment_variables import MLFLOW_LOGGING_LEVEL - from mlflow.utils.logging_utils import SuppressLogFilter import logging.config - import rich.logging import os + from mlflow.environment_variables import MLFLOW_LOGGING_LEVEL + from mlflow.utils.logging_utils import SuppressLogFilter + if 'MLFLOW_SCORING_SERVER_REQUEST_TIMEOUT' not in os.environ: # probably not running in mlflow server, so no need to configure our mlflow loggers return @@ -114,4 +116,4 @@ def _configure_mlflow_loggers(): ) -__all__ = ['set_project', 'setup_mlflow_environment', 'ensure_mlflow_configured', 'BaseDatamintModel', 'DatamintModel'] +__all__ = ['BaseDatamintModel', 'DatamintModel', 'ensure_mlflow_configured', 'set_project', 'setup_mlflow_environment'] diff --git a/datamint/mlflow/artifact/__init__.py b/datamint/mlflow/artifact/__init__.py index 6c0799a3..f39d1f68 100644 --- a/datamint/mlflow/artifact/__init__.py +++ b/datamint/mlflow/artifact/__init__.py @@ -1 +1 @@ -from .datamint_artifacts_repo import DatamintArtifactsRepository \ No newline at end of file +from .datamint_artifacts_repo import DatamintArtifactsRepository as DatamintArtifactsRepository \ No newline at end of file diff --git a/datamint/mlflow/data/__init__.py b/datamint/mlflow/data/__init__.py index 5655137f..3208354c 100644 --- a/datamint/mlflow/data/__init__.py +++ b/datamint/mlflow/data/__init__.py @@ -1,3 +1,3 @@ -from .datamint_dataset import DatamintMLflowDataset, DatamintDatasetSource +from .datamint_dataset import DatamintDatasetSource, DatamintMLflowDataset -__all__ = ["DatamintMLflowDataset", "DatamintDatasetSource"] +__all__ = ["DatamintDatasetSource", "DatamintMLflowDataset"] diff --git a/datamint/mlflow/data/datamint_dataset.py b/datamint/mlflow/data/datamint_dataset.py index c9d8008c..d84a4f4f 100644 --- a/datamint/mlflow/data/datamint_dataset.py +++ b/datamint/mlflow/data/datamint_dataset.py @@ -3,14 +3,14 @@ import hashlib import json -from typing import Any -from collections.abc import Sequence import logging +from collections.abc import Sequence +from typing import Any from mlflow.data.dataset import Dataset from mlflow.data.dataset_source import DatasetSource -from datamint.entities.resource import Resource +from datamint.entities.resource import Resource _LOGGER = logging.getLogger(__name__) diff --git a/datamint/mlflow/env_utils.py b/datamint/mlflow/env_utils.py index bef8c3ed..1c0c3dff 100644 --- a/datamint/mlflow/env_utils.py +++ b/datamint/mlflow/env_utils.py @@ -3,12 +3,12 @@ based on Datamint configuration. """ -import os import logging -from urllib.parse import urlparse -from datamint import configs +import os import sys +from urllib.parse import urlparse +from datamint import configs _LOGGER = logging.getLogger(__name__) @@ -92,7 +92,7 @@ def setup_mlflow_environment(overwrite: bool = False, if 'lightning.pytorch.loggers' in sys.modules: # import lightning.pytorch.loggers # importlib.reload(lightning.pytorch.loggers) - from lightning.pytorch.loggers import MLFlowLogger + from lightning.pytorch.loggers import MLFlowLogger # 1. Convert the immutable defaults tuple to a mutable list current_defaults = list(MLFlowLogger.__init__.__defaults__) diff --git a/datamint/mlflow/env_vars.py b/datamint/mlflow/env_vars.py index 6a47d97a..dfa37877 100644 --- a/datamint/mlflow/env_vars.py +++ b/datamint/mlflow/env_vars.py @@ -1,5 +1,6 @@ from enum import Enum + class EnvVars(Enum): DATAMINT_PROJECT_ID = "DATAMINT_PROJECT_ID" DATAMINT_PROJECT_NAME = "DATAMINT_PROJECT_NAME" diff --git a/datamint/mlflow/flavors/__init__.py b/datamint/mlflow/flavors/__init__.py index dce6ea42..b555e723 100644 --- a/datamint/mlflow/flavors/__init__.py +++ b/datamint/mlflow/flavors/__init__.py @@ -3,22 +3,27 @@ """ from .datamint_flavor import ( - save_model, - log_model, - load_model, _load_pyfunc, + load_model, + log_model, + save_model, ) from .task_type import TaskType -from .validation import validate_model, ValidationReport, ValidationIssue, ModelValidationError +from .validation import ( + ModelValidationError, + ValidationIssue, + ValidationReport, + validate_model, +) __all__ = [ - "save_model", - "log_model", - "load_model", - "_load_pyfunc", + "ModelValidationError", "TaskType", - "validate_model", - "ValidationReport", "ValidationIssue", - "ModelValidationError", + "ValidationReport", + "_load_pyfunc", + "load_model", + "log_model", + "save_model", + "validate_model", ] diff --git a/datamint/mlflow/flavors/datamint_flavor.py b/datamint/mlflow/flavors/datamint_flavor.py index ac9e8790..707314b3 100644 --- a/datamint/mlflow/flavors/datamint_flavor.py +++ b/datamint/mlflow/flavors/datamint_flavor.py @@ -1,20 +1,23 @@ import logging +import tempfile +from collections.abc import Sequence +from dataclasses import asdict from pathlib import Path +from typing import Any + import mlflow +import torch +from mlflow import pyfunc from mlflow.models import Model, ModelInputExample, ModelSignature +from mlflow.pytorch import pickle_module as mlflow_pytorch_pickle_module +from packaging.requirements import Requirement + import datamint import datamint.mlflow.flavors -from mlflow import pyfunc +from datamint.entities.annotations.annotation_spec import AnnotationSpec + from .model import BaseDatamintModel, DatamintModel, _DatamintModelWrapper from .task_type import TaskType -from datamint.entities.annotations.annotation_spec import AnnotationSpec -from collections.abc import Sequence -from dataclasses import asdict -from packaging.requirements import Requirement -from typing import Any -import torch -import tempfile -from mlflow.pytorch import pickle_module as mlflow_pytorch_pickle_module logger = logging.getLogger(__name__) @@ -30,18 +33,18 @@ def _process_input_example(input_example: ModelInputExample | None) -> tuple[Mod if input_example is None: import datetime - input_resource = dict( - id='model_id', - storage='DicomResource', - filename='file.dcm', - location='private/location', - mimetype='application/dicom', - size=14724562, - status='inbox', - created_at=datetime.datetime.now().isoformat(), - created_by='user@mail.com', - modality='CT' - ) + input_resource = { + 'id': 'model_id', + 'storage': 'DicomResource', + 'filename': 'file.dcm', + 'location': 'private/location', + 'mimetype': 'application/dicom', + 'size': 14724562, + 'status': 'inbox', + 'created_at': datetime.datetime.now().isoformat(), + 'created_by': 'user@mail.com', + 'modality': 'CT' + } return [input_resource], datamint_params if not isinstance(input_example, tuple): return (input_example, datamint_params) @@ -90,9 +93,9 @@ def _build_datamint_wheel(source_dir: Path) -> tuple[str, Path]: resolve to system paths outside the project root, causing poetry-core to crash with a ValueError. Building from a clean copy avoids this. """ - import sys - import subprocess import shutil + import subprocess + import sys import tempfile as _tmp _IGNORE = shutil.ignore_patterns( @@ -295,7 +298,7 @@ def save_model(datamint_model: BaseDatamintModel, # DatamintLightningModule is an nn.Module itself, so its CUDA weights are # embedded directly in the cloudpickle. Move to CPU before serialization # so the pickle is device-agnostic and loads on CPU-only containers. - import torch.nn as nn + from torch import nn _underlying = datamint_model.another_model if isinstance(datamint_model, _DatamintModelWrapper) else datamint_model if isinstance(_underlying, nn.Module): _underlying.cpu() diff --git a/datamint/mlflow/flavors/model.py b/datamint/mlflow/flavors/model.py index c801c1e7..a9efb8fe 100644 --- a/datamint/mlflow/flavors/model.py +++ b/datamint/mlflow/flavors/model.py @@ -1,24 +1,30 @@ +import logging from abc import ABC from collections.abc import Sequence from dataclasses import dataclass from functools import cached_property -import logging from typing import Any, ClassVar, TypeAlias +import torch +from medimgkit import ViewPlane from mlflow.environment_variables import MLFLOW_DEFAULT_PREDICTION_DEVICE from mlflow.pyfunc import PyFuncModel from mlflow.pyfunc.model import PythonModel, PythonModelContext -from medimgkit import ViewPlane -import torch from typing_extensions import override -from datamint.entities.annotations import Annotation, ImageSegmentation, ImageClassification +from datamint.entities.annotations import ( + Annotation, + ImageClassification, + ImageSegmentation, +) from datamint.entities.annotations.annotation_spec import AnnotationSpec from datamint.entities.resource import BaseResource from datamint.entities.resources.volume_resource import VolumeResource from datamint.entities.sliced_resource import SlicedVolumeResource from datamint.entities.sliced_video_resource import SlicedVideoResource -from datamint.mlflow.flavors.model_loader import LINKED_MODELS_DIR as DEFAULT_LINKED_MODELS_DIR +from datamint.mlflow.flavors.model_loader import ( + LINKED_MODELS_DIR as DEFAULT_LINKED_MODELS_DIR, +) from datamint.mlflow.flavors.model_loader import LinkedModelLoader from datamint.mlflow.flavors.prediction_router import PredictionRouter, bridge_mode from datamint.mlflow.flavors.task_type import TaskType diff --git a/datamint/mlflow/flavors/prediction_router.py b/datamint/mlflow/flavors/prediction_router.py index 2e4be7c4..897bd2ba 100644 --- a/datamint/mlflow/flavors/prediction_router.py +++ b/datamint/mlflow/flavors/prediction_router.py @@ -10,7 +10,8 @@ import logging from collections.abc import Callable from dataclasses import dataclass -from typing import Any +from typing import Any, ClassVar + from .prediction_modes import PredictionMode _LOGGER = logging.getLogger(__name__) @@ -69,11 +70,11 @@ class PredictionRouter: """ _RESERVED_PARAMS = frozenset({"mode", "confidence_threshold"}) - _DELEGATED_MODE_MAP = { + _DELEGATED_MODE_MAP: ClassVar[dict[PredictionMode, PredictionMode]] = { PredictionMode.IMAGE: PredictionMode.SLICE, } # prerequisite → bridge modes that become available when prerequisite is registered - _BRIDGE_MODE_PREREQS: dict[PredictionMode, PredictionMode] = { + _BRIDGE_MODE_PREREQS: ClassVar[dict[PredictionMode, PredictionMode]] = { PredictionMode.VOLUME: PredictionMode.SLICE, PredictionMode.FRAME: PredictionMode.IMAGE, PredictionMode.ALL_FRAMES: PredictionMode.IMAGE, diff --git a/datamint/mlflow/flavors/validation.py b/datamint/mlflow/flavors/validation.py index fc17de63..2dc23baa 100644 --- a/datamint/mlflow/flavors/validation.py +++ b/datamint/mlflow/flavors/validation.py @@ -7,9 +7,10 @@ if TYPE_CHECKING: from mlflow.models import ModelSignature - from datamint.mlflow.flavors.model import BaseDatamintModel + from datamint.dataset.base import DatamintBaseDataset from datamint.entities.resource import BaseResource + from datamint.mlflow.flavors.model import BaseDatamintModel _LOGGER = logging.getLogger(__name__) @@ -188,7 +189,7 @@ def _check_inference( f'got: {", ".join(sorted(wrong))}', 'error')) else: issues.append(ValidationIssue('task_type_consistency', True, - f'annotation types consistent with task_type')) + 'annotation types consistent with task_type')) return issues, signature diff --git a/datamint/mlflow/lightning/callbacks/__init__.py b/datamint/mlflow/lightning/callbacks/__init__.py index 4d5d971f..00cadc87 100644 --- a/datamint/mlflow/lightning/callbacks/__init__.py +++ b/datamint/mlflow/lightning/callbacks/__init__.py @@ -1,5 +1,5 @@ from .modelcheckpoint import ( - MLFlowModelCheckpoint, - MLFlowPyTorchModelCheckpoint, - MLFlowDatamintModelCheckpoint, -) \ No newline at end of file + MLFlowDatamintModelCheckpoint as MLFlowDatamintModelCheckpoint, + MLFlowModelCheckpoint as MLFlowModelCheckpoint, + MLFlowPyTorchModelCheckpoint as MLFlowPyTorchModelCheckpoint, +) diff --git a/datamint/mlflow/lightning/callbacks/modelcheckpoint.py b/datamint/mlflow/lightning/callbacks/modelcheckpoint.py index 64bcb0b0..dd0fe7e1 100644 --- a/datamint/mlflow/lightning/callbacks/modelcheckpoint.py +++ b/datamint/mlflow/lightning/callbacks/modelcheckpoint.py @@ -1,32 +1,35 @@ +import copy +import hashlib +import inspect +import json +import logging from collections.abc import Mapping -from lightning.pytorch.callbacks import ModelCheckpoint +from concurrent.futures import Future, ThreadPoolExecutor from pathlib import Path +from typing import TYPE_CHECKING, Any, Literal from weakref import proxy -from mlflow.store.artifact.artifact_repository_registry import get_artifact_repository -from typing import Literal, Any, TYPE_CHECKING -from typing_extensions import override -import inspect -from torch import nn + import lightning.pytorch as L -from datamint.mlflow.models import log_model_metadata, _get_MLFlowLogger -from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG -from datamint.mlflow.env_utils import ensure_mlflow_configured -import mlflow.models -import mlflow.exceptions -import mlflow.pytorch import mlflow.data.dataset import mlflow.entities.dataset -import copy -import logging -import json -import hashlib -from concurrent.futures import ThreadPoolExecutor, Future +import mlflow.exceptions +import mlflow.models +import mlflow.pytorch +from lightning.pytorch.callbacks import ModelCheckpoint from lightning.pytorch.loggers import MLFlowLogger +from mlflow.store.artifact.artifact_repository_registry import get_artifact_repository +from torch import nn +from typing_extensions import override + +from datamint.mlflow.env_utils import ensure_mlflow_configured +from datamint.mlflow.models import _get_MLFlowLogger, log_model_metadata +from datamint.mlflow.models.tags import DATAMINT_LOGGED_MODEL_ID_TAG if TYPE_CHECKING: - from datamint.mlflow.flavors.model import BaseDatamintModel from mlflow.models.model import ModelInfo + from datamint.mlflow.flavors.model import BaseDatamintModel + _LOGGER = logging.getLogger(__name__) @@ -138,7 +141,7 @@ def _inject_model_id(self, model: 'nn.Module | L.LightningModule | BaseDatamintM if hasattr(model, 'set_mlflow_model_id'): model.set_mlflow_model_id(self._last_model_id) elif hasattr(model, 'mlflow_model_id'): - setattr(model, 'mlflow_model_id', self._last_model_id) + model.mlflow_model_id = self._last_model_id def _prepare_loggable_model(self, model: nn.Module) -> nn.Module: """Prepare a model for MLflow logging, potentially creating a CPU copy. @@ -389,7 +392,6 @@ def _wrap_forward(self, pl_module: nn.Module) -> None: Override in subclasses to customize signature inference. """ - pass def on_fit_start(self, trainer: L.Trainer, pl_module: L.LightningModule) -> None: super().on_fit_start(trainer, pl_module) @@ -452,7 +454,7 @@ def _restore_model_uri(self, trainer: L.Trainer) -> None: _LOGGER.warning("MLFlowLogger has no run_id. Cannot restore model URI.") return if logger.run_id not in str(trainer.ckpt_path): - _LOGGER.warning(f"Run ID mismatch between checkpoint path and MLFlowLogger." + + _LOGGER.warning("Run ID mismatch between checkpoint path and MLFlowLogger." + " Check `run_id` parameter in MLFlowLogger.") return model_name = Path(trainer.ckpt_path).stem[:256] diff --git a/datamint/mlflow/models/__init__.py b/datamint/mlflow/models/__init__.py index 43cbcd48..7a21c453 100644 --- a/datamint/mlflow/models/__init__.py +++ b/datamint/mlflow/models/__init__.py @@ -1,11 +1,12 @@ -import logging import json -import lightning as L -from lightning.pytorch.loggers import MLFlowLogger -import mlflow +import logging import os from tempfile import TemporaryDirectory +import lightning as L +import mlflow +from lightning.pytorch.loggers import MLFlowLogger + _LOGGER = logging.getLogger(__name__) diff --git a/datamint/mlflow/models/datamint_model_store.py b/datamint/mlflow/models/datamint_model_store.py index b6685a0a..c00835c0 100644 --- a/datamint/mlflow/models/datamint_model_store.py +++ b/datamint/mlflow/models/datamint_model_store.py @@ -1,12 +1,6 @@ -from mlflow.store.model_registry.rest_store import RestStore -from datamint.mlflow.store_utils import _inject_project_id_into_body -from datamint.mlflow.tracking.fluent import get_active_project_id -from mlflow.exceptions import MlflowException -from mlflow.utils.proto_json_utils import message_to_json -from typing_extensions import override -from mlflow.entities.model_registry import ModelVersion, RegisteredModel from functools import partial +from mlflow.entities.model_registry import ModelVersion, RegisteredModel from mlflow.protos.model_registry_pb2 import ( CreateModelVersion, CreateRegisteredModel, @@ -20,7 +14,6 @@ GetModelVersionByAlias, GetModelVersionDownloadUri, GetRegisteredModel, - ModelRegistryService, RenameRegisteredModel, SearchModelVersions, SearchRegisteredModels, @@ -31,6 +24,12 @@ UpdateModelVersion, UpdateRegisteredModel, ) +from mlflow.store.model_registry.rest_store import RestStore +from mlflow.utils.proto_json_utils import message_to_json +from typing_extensions import override + +from datamint.mlflow.store_utils import _inject_project_id_into_body +from datamint.mlflow.tracking.fluent import get_active_project_id class DatamintModelRegistryStore(RestStore): @@ -44,9 +43,10 @@ class DatamintModelRegistryStore(RestStore): def __init__(self, store_uri: str, artifact_uri=None, force_valid=True): # Ensure MLflow environment is configured when store is initialized - from datamint.mlflow.env_utils import setup_mlflow_environment from mlflow.utils.credentials import get_default_host_creds + from datamint.mlflow.env_utils import setup_mlflow_environment + setup_mlflow_environment() if store_uri.startswith('datamint://') or 'datamint.io' in store_uri or force_valid: diff --git a/datamint/mlflow/store_utils.py b/datamint/mlflow/store_utils.py index a9b1c866..d2fc1672 100644 --- a/datamint/mlflow/store_utils.py +++ b/datamint/mlflow/store_utils.py @@ -1,7 +1,9 @@ -from mlflow.protos.databricks_pb2 import INVALID_PARAMETER_VALUE -from mlflow.exceptions import MlflowException import json +from mlflow.exceptions import MlflowException +from mlflow.protos.databricks_pb2 import INVALID_PARAMETER_VALUE + + def _resolve_project_id(project_id: str | None) -> str: """ Resolve the project ID from the provided value or active context. diff --git a/datamint/mlflow/tracking/datamint_store.py b/datamint/mlflow/tracking/datamint_store.py index ea306146..f4e8f4b0 100644 --- a/datamint/mlflow/tracking/datamint_store.py +++ b/datamint/mlflow/tracking/datamint_store.py @@ -1,9 +1,14 @@ -from mlflow.store.tracking.rest_store import RestStore +from functools import partial + from mlflow.exceptions import MlflowException +from mlflow.store.tracking.rest_store import RestStore from mlflow.utils.proto_json_utils import message_to_json -from functools import partial from typing_extensions import override -from datamint.mlflow.store_utils import _resolve_project_id, _inject_project_id_into_body + +from datamint.mlflow.store_utils import ( + _inject_project_id_into_body, + _resolve_project_id, +) class DatamintStore(RestStore): @@ -14,8 +19,9 @@ class DatamintStore(RestStore): def __init__(self, store_uri: str, artifact_uri=None, force_valid=True): # Ensure MLflow environment is configured when store is initialized - from datamint.mlflow.env_utils import setup_mlflow_environment from mlflow.utils.credentials import get_default_host_creds + + from datamint.mlflow.env_utils import setup_mlflow_environment setup_mlflow_environment() if store_uri.startswith('datamint://') or 'datamint.io' in store_uri or force_valid: @@ -46,9 +52,9 @@ def create_experiment(self, name, artifact_location=None, tags=None, project_id: @override def get_experiment_by_name(self, experiment_name, project_id: str | None = None): - from mlflow.protos.service_pb2 import GetExperimentByName from mlflow.entities import Experiment from mlflow.protos import databricks_pb2 + from mlflow.protos.service_pb2 import GetExperimentByName if self.invalid: return super().get_experiment_by_name(experiment_name) diff --git a/datamint/mlflow/tracking/default_experiment.py b/datamint/mlflow/tracking/default_experiment.py index e40f4efb..700f2148 100644 --- a/datamint/mlflow/tracking/default_experiment.py +++ b/datamint/mlflow/tracking/default_experiment.py @@ -1,6 +1,9 @@ -import sys import os -from mlflow.tracking.default_experiment.abstract_context import DefaultExperimentProvider +import sys + +from mlflow.tracking.default_experiment.abstract_context import ( + DefaultExperimentProvider, +) class DatamintExperimentProvider(DefaultExperimentProvider): diff --git a/datamint/mlflow/tracking/fluent.py b/datamint/mlflow/tracking/fluent.py index 643a4898..66e08aa9 100644 --- a/datamint/mlflow/tracking/fluent.py +++ b/datamint/mlflow/tracking/fluent.py @@ -1,11 +1,12 @@ -from typing import TYPE_CHECKING -import threading import logging +import os +import threading +from typing import TYPE_CHECKING + from datamint import Api from datamint.exceptions import ItemNotFoundError -import os -from datamint.mlflow.env_vars import EnvVars from datamint.mlflow.env_utils import ensure_mlflow_configured +from datamint.mlflow.env_vars import EnvVars if TYPE_CHECKING: from datamint.entities.project import Project diff --git a/datamint/types.py b/datamint/types.py index b426487c..15f4dbbf 100644 --- a/datamint/types.py +++ b/datamint/types.py @@ -1,10 +1,10 @@ -from typing import Literal, TypeAlias, TYPE_CHECKING, Union +from typing import TYPE_CHECKING, Literal, TypeAlias, Union if TYPE_CHECKING: - import pydicom.dataset - from PIL import Image import cv2 + import pydicom.dataset from nibabel.filebasedimages import FileBasedImage as nib_FileBasedImage + from PIL import Image # Type alias for imaging formats ImagingData: TypeAlias = ( diff --git a/datamint/utils/annotation_agreement.py b/datamint/utils/annotation_agreement.py index 71c341f1..0713d52b 100644 --- a/datamint/utils/annotation_agreement.py +++ b/datamint/utils/annotation_agreement.py @@ -7,9 +7,10 @@ from __future__ import annotations +from collections.abc import Sequence from dataclasses import dataclass from itertools import combinations -from typing import Literal, Sequence +from typing import Literal import numpy as np import pandas as pd diff --git a/datamint/utils/collection_utils.py b/datamint/utils/collection_utils.py index e1a79e04..b12ceb8f 100644 --- a/datamint/utils/collection_utils.py +++ b/datamint/utils/collection_utils.py @@ -16,7 +16,7 @@ class ChainedSequence(Sequence[T]): methods like __iter__, __contains__, __reversed__, count(), and index(). """ - __slots__ = ('_sequences', '_cumulative_lengths') + __slots__ = ('_cumulative_lengths', '_sequences') def __init__(self, *sequences: Sequence[T]) -> None: self._sequences = sequences diff --git a/datamint/utils/logging_utils.py b/datamint/utils/logging_utils.py index 4fa1d5c4..9488601c 100644 --- a/datamint/utils/logging_utils.py +++ b/datamint/utils/logging_utils.py @@ -1,16 +1,16 @@ -from rich.theme import Theme -from logging import Logger, DEBUG, INFO, WARNING, ERROR, CRITICAL -from rich.console import Console -import platform -import os +import datetime +import importlib import logging import logging.config -from rich.console import ConsoleRenderable +import os +import platform +from logging import CRITICAL, DEBUG, ERROR, INFO, WARNING + +import yaml +from rich.console import Console, ConsoleRenderable from rich.logging import RichHandler +from rich.theme import Theme from rich.traceback import Traceback -import yaml -import importlib -import datetime _LOGGER = logging.getLogger(__name__) @@ -53,7 +53,7 @@ def load_cmdline_logging_config(): # try loading the developer's logging config with open('logging_dev.yaml', 'r') as f: config = yaml.safe_load(f) - except: + except Exception: with importlib.resources.open_text('datamint', 'logging.yaml') as f: config = yaml.safe_load(f.read()) diff --git a/datamint/utils/nifti_utils.py b/datamint/utils/nifti_utils.py index 31a932c3..d8655653 100644 --- a/datamint/utils/nifti_utils.py +++ b/datamint/utils/nifti_utils.py @@ -1,6 +1,7 @@ -from nibabel.nifti1 import Nifti1Header, Nifti1Image from collections.abc import Mapping + import numpy as np +from nibabel.nifti1 import Nifti1Header, Nifti1Image def _get_metadata_value(metadata: Mapping[str, object], *keys: str) -> object | None: @@ -123,7 +124,7 @@ def _build_nifti_header_from_metadata(metadata: Mapping[str, object], def metadata_to_nifti_obj(metadata: Mapping[str, object], dataobj: np.ndarray | None = None, *, - fill_value: int | float = 0) -> Nifti1Image: + fill_value: float = 0) -> Nifti1Image: """Construct a ``Nifti1Image`` from a metadata mapping and optional data. nibabel provides the building blocks for this via ``Nifti1Header`` and diff --git a/datamint/utils/torchmetrics.py b/datamint/utils/torchmetrics.py index a32833ba..9ac5990d 100644 --- a/datamint/utils/torchmetrics.py +++ b/datamint/utils/torchmetrics.py @@ -1,8 +1,7 @@ import torch -from torchmetrics.classification import Recall, Precision, F1Score, Specificity -from torchmetrics.wrappers.abstract import WrapperMetric import torchmetrics from torch import Tensor +from torchmetrics.wrappers.abstract import WrapperMetric class SegmentationToClassificationWrapper(WrapperMetric): @@ -78,7 +77,7 @@ class CombinedLoss(torch.nn.Module): def __init__(self, *losses: torch.nn.Module | tuple[torch.nn.Module, float]): super().__init__() - parsed = [(l[0], l[1]) if isinstance(l, tuple) else (l, 1.0) for l in losses] + parsed = [(loss[0], loss[1]) if isinstance(loss, tuple) else (loss, 1.0) for loss in losses] self.losses, self.weights = zip(*parsed) def forward(self, preds: Tensor, target: Tensor) -> Tensor: diff --git a/datamint/utils/visualization.py b/datamint/utils/visualization.py index c6d5ed05..78980519 100644 --- a/datamint/utils/visualization.py +++ b/datamint/utils/visualization.py @@ -1,13 +1,14 @@ +import colorsys +import logging +from collections.abc import Sequence + import matplotlib.pyplot as plt import numpy as np -from torchvision.transforms import functional as F -from torch import Tensor -import torchvision.utils import torch -import colorsys -from collections.abc import Sequence +import torchvision.utils from matplotlib.axes import Axes -import logging +from torch import Tensor +from torchvision.transforms import functional as F _LOGGER = logging.getLogger(__name__) @@ -41,9 +42,9 @@ def show(imgs: Sequence[Tensor | np.ndarray] | Tensor | np.ndarray, if ax is None: if figsize is not None: - fig, axs = plt.subplots(ncols=len(imgs), squeeze=False, figsize=figsize) + _fig, axs = plt.subplots(ncols=len(imgs), squeeze=False, figsize=figsize) else: - fig, axs = plt.subplots(ncols=len(imgs), squeeze=False) + _fig, axs = plt.subplots(ncols=len(imgs), squeeze=False) axs = axs[0] else: if isinstance(ax, Axes): diff --git a/notebooks/01_getting_started/01_upload_data.ipynb b/notebooks/01_getting_started/01_upload_data.ipynb index deb58ac7..203815dc 100644 --- a/notebooks/01_getting_started/01_upload_data.ipynb +++ b/notebooks/01_getting_started/01_upload_data.ipynb @@ -48,8 +48,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint import Api\n", "from pathlib import Path\n", + "\n", + "from datamint import Api\n", + "\n", "# Creates a connection with the server.\n", "# Don't forget to run `datamint config` in a terminal, if you haven't already.\n", "# Or use api_key parameter in Api()\n", @@ -326,7 +328,34 @@ "execution_count": null, "metadata": {}, "outputs": [], - "source": "# Get some resources to add to a project\ntutorial_resources = list(api.resources.get_list(\n tags=['tutorial'],\n status='inbox'\n))\n\nif tutorial_resources:\n resource_ids_for_project = [r.id for r in tutorial_resources[:3]] # Take first 3 resources\n\n # Create a new project (exists_ok returns the existing project instead of raising)\n project = api.projects.create(\n name=\"Tutorial Project\",\n description=\"A project created for demonstration purposes\",\n resource_ids=resource_ids_for_project,\n exists_ok=True\n )\n\n print(f\"Created project: {project.name} (ID: {project.id})\")\n\n # List all projects\n all_projects = api.projects.get_list()\n print(f\"\\nAll projects ({len(all_projects)}):\")\n for proj in all_projects:\n print(f\" - {proj.name} (ID: {proj.id})\")\nelse:\n print(\"No tutorial resources found to add to project\")" + "source": [ + "# Get some resources to add to a project\n", + "tutorial_resources = list(api.resources.get_list(\n", + " tags=['tutorial'],\n", + " status='inbox'\n", + "))\n", + "\n", + "if tutorial_resources:\n", + " resource_ids_for_project = [r.id for r in tutorial_resources[:3]] # Take first 3 resources\n", + "\n", + " # Create a new project (exists_ok returns the existing project instead of raising)\n", + " project = api.projects.create(\n", + " name=\"Tutorial Project\",\n", + " description=\"A project created for demonstration purposes\",\n", + " resource_ids=resource_ids_for_project,\n", + " exists_ok=True\n", + " )\n", + "\n", + " print(f\"Created project: {project.name} (ID: {project.id})\")\n", + "\n", + " # List all projects\n", + " all_projects = api.projects.get_list()\n", + " print(f\"\\nAll projects ({len(all_projects)}):\")\n", + " for proj in all_projects:\n", + " print(f\" - {proj.name} (ID: {proj.id})\")\n", + "else:\n", + " print(\"No tutorial resources found to add to project\")" + ] }, { "cell_type": "markdown", diff --git a/notebooks/01_getting_started/02_explore_data.ipynb b/notebooks/01_getting_started/02_explore_data.ipynb index 679c3204..e2a9b777 100644 --- a/notebooks/01_getting_started/02_explore_data.ipynb +++ b/notebooks/01_getting_started/02_explore_data.ipynb @@ -35,11 +35,13 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint import Api\n", + "import os\n", + "\n", "import matplotlib.pyplot as plt\n", - "from PIL import Image\n", "import numpy as np\n", - "import os\n", + "from PIL import Image\n", + "\n", + "from datamint import Api\n", "\n", "# Initialize the API client\n", "# Make sure you have setup your api key\n", @@ -325,7 +327,6 @@ "metadata": {}, "outputs": [], "source": [ - "from io import BytesIO\n", "\n", "# Filter for segmentation annotations\n", "segmentation_annotations = [ann for ann in annotations if ann.type == 'segmentation']\n", @@ -373,7 +374,7 @@ " axes[0].axis('off')\n", " \n", " # Segmentation\n", - " axes[1].imshow(seg_array, cmap='jet', alpha=0.7)\n", + " axes[1].imshow(seg_data, cmap='jet', alpha=0.7)\n", " axes[1].set_title(f\"Segmentation: {seg_annotation.identifier}\")\n", " axes[1].axis('off')\n", " \n", @@ -445,7 +446,34 @@ "id": "175ee2f9", "metadata": {}, "outputs": [], - "source": "from datetime import date, timedelta\n\n# Filter by annotation type\ncategory_annotations = api.annotations.get_list(\n resource=selected_resource,\n annotation_type='category'\n)\nprint(f\"Category annotations: {len(category_annotations)}\")\n\n# Filter by date range\ndate_to = date.today()\ndate_from = date_to - timedelta(days=30) # Last 30 days\n\nrecent_annotations = api.annotations.get_list(\n resource=selected_resource,\n from_date=date_from,\n to_date=date_to\n)\nprint(f\"Annotations from last 30 days: {len(recent_annotations)}\")\n\n# Filter by status\npublished_annotations = api.annotations.get_list(\n resource=selected_resource,\n status='published'\n)\nprint(f\"Published annotations: {len(published_annotations)}\")" + "source": [ + "from datetime import date, timedelta\n", + "\n", + "# Filter by annotation type\n", + "category_annotations = api.annotations.get_list(\n", + " resource=selected_resource,\n", + " annotation_type='category'\n", + ")\n", + "print(f\"Category annotations: {len(category_annotations)}\")\n", + "\n", + "# Filter by date range\n", + "date_to = date.today()\n", + "date_from = date_to - timedelta(days=30) # Last 30 days\n", + "\n", + "recent_annotations = api.annotations.get_list(\n", + " resource=selected_resource,\n", + " from_date=date_from,\n", + " to_date=date_to\n", + ")\n", + "print(f\"Annotations from last 30 days: {len(recent_annotations)}\")\n", + "\n", + "# Filter by status\n", + "published_annotations = api.annotations.get_list(\n", + " resource=selected_resource,\n", + " status='published'\n", + ")\n", + "print(f\"Published annotations: {len(published_annotations)}\")" + ] }, { "cell_type": "markdown", @@ -489,4 +517,4 @@ }, "nbformat": 4, "nbformat_minor": 5 -} \ No newline at end of file +} diff --git a/notebooks/03_datasets/01_project_scoped_splits.ipynb b/notebooks/03_datasets/01_project_scoped_splits.ipynb index 23c87bc1..174ce4f4 100644 --- a/notebooks/03_datasets/01_project_scoped_splits.ipynb +++ b/notebooks/03_datasets/01_project_scoped_splits.ipynb @@ -99,6 +99,7 @@ "source": [ "from random import Random\n", "\n", + "\n", "def partition_ids(resource_ids: list[str], ratios: dict[str, float], seed: int = 42) -> dict[str, list[str]]:\n", " total = sum(ratios.values())\n", " if abs(total - 1.0) > 1e-6:\n", diff --git a/notebooks/03_datasets/02_patient_wise_splits.ipynb b/notebooks/03_datasets/02_patient_wise_splits.ipynb index 133d9c07..11de312c 100644 --- a/notebooks/03_datasets/02_patient_wise_splits.ipynb +++ b/notebooks/03_datasets/02_patient_wise_splits.ipynb @@ -61,9 +61,10 @@ } ], "source": [ + "from pathlib import Path\n", + "\n", "import pydicom\n", "import pydicom.data\n", - "from pathlib import Path\n", "\n", "from datamint import Api\n", "from datamint.dataset import ImageDataset\n", diff --git a/notebooks/03_datasets/03_build_dataset.ipynb b/notebooks/03_datasets/03_build_dataset.ipynb index 2f881502..971c27c5 100644 --- a/notebooks/03_datasets/03_build_dataset.ipynb +++ b/notebooks/03_datasets/03_build_dataset.ipynb @@ -52,8 +52,8 @@ } ], "source": [ - "import numpy as np\n", "import matplotlib.pyplot as plt\n", + "\n", "from datamint.dataset import build_dataset" ] }, diff --git a/notebooks/03_datasets/04_volume_dataset.ipynb b/notebooks/03_datasets/04_volume_dataset.ipynb index 48022b7b..5d760aad 100644 --- a/notebooks/03_datasets/04_volume_dataset.ipynb +++ b/notebooks/03_datasets/04_volume_dataset.ipynb @@ -46,9 +46,8 @@ "metadata": {}, "outputs": [], "source": [ - "import numpy as np\n", "import matplotlib.pyplot as plt\n", - "from datamint.dataset import build_dataset\n", + "\n", "from datamint.dataset import VolumeDataset" ] }, diff --git a/notebooks/04_experiment_tracking/01_mlflow_manual_logging.ipynb b/notebooks/04_experiment_tracking/01_mlflow_manual_logging.ipynb index bba6f43a..133907d3 100644 --- a/notebooks/04_experiment_tracking/01_mlflow_manual_logging.ipynb +++ b/notebooks/04_experiment_tracking/01_mlflow_manual_logging.ipynb @@ -18,6 +18,7 @@ "# Core MLflow imports\n", "import mlflow\n", "import mlflow.pytorch\n", + "\n", "from datamint import Api\n", "from datamint.mlflow import set_project" ] diff --git a/notebooks/05_deployment/01_deploy_registered_model.ipynb b/notebooks/05_deployment/01_deploy_registered_model.ipynb index e1e45892..fd165f2f 100644 --- a/notebooks/05_deployment/01_deploy_registered_model.ipynb +++ b/notebooks/05_deployment/01_deploy_registered_model.ipynb @@ -41,11 +41,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint import Api\n", - "import datamint.mlflow\n", - "import mlflow\n", "import time\n", "\n", + "from datamint import Api\n", + "\n", "# Initialize the API client\n", "# The API key can be set via environment variable DATAMINT_API_KEY\n", "# or passed directly to the Api constructor\n", @@ -70,9 +69,10 @@ "outputs": [], "source": [ "# get all versions of all registered models\n", - "from mlflow import MlflowClient\n", "from datetime import datetime\n", "\n", + "from mlflow import MlflowClient\n", + "\n", "client = MlflowClient()\n", "all_registered_models = client.search_registered_models()\n", "for rm in all_registered_models:\n", @@ -113,7 +113,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")\n", "print(f\"Model: {job.model_name}\")" @@ -410,7 +410,7 @@ " final_status = wait_for_job_completion(job.id, poll_interval=5)\n", " \n", " if final_status and final_status.status == 'completed':\n", - " print(f\"\\n3. Deployment completed successfully!\")\n", + " print(\"\\n3. Deployment completed successfully!\")\n", " print(f\" Image: {final_status.image_name}:{final_status.image_tag}\")\n", " \n", " # Verify image exists\n", @@ -422,7 +422,7 @@ " \n", " return final_status\n", " else:\n", - " print(f\"\\n3. Deployment failed or timed out\")\n", + " print(\"\\n3. Deployment failed or timed out\")\n", " if final_status and final_status.error_message:\n", " print(f\" Error: {final_status.error_message}\")\n", " return None\n", diff --git a/notebooks/05_deployment/02_deploy_external_model.ipynb b/notebooks/05_deployment/02_deploy_external_model.ipynb index 9a8c0e9a..5c8dc216 100644 --- a/notebooks/05_deployment/02_deploy_external_model.ipynb +++ b/notebooks/05_deployment/02_deploy_external_model.ipynb @@ -100,9 +100,8 @@ "metadata": {}, "outputs": [], "source": [ - "import torch\n", - "import torch.nn as nn\n", "import segmentation_models_pytorch as smp\n", + "import torch\n", "\n", "CLASS_NAMES = ['lesion'] # foreground class names; one per output channel\n", "IMAGE_SIZE = 256\n", @@ -173,15 +172,14 @@ "metadata": {}, "outputs": [], "source": [ - "import numpy as np\n", "import albumentations as A\n", + "import cv2\n", + "import numpy as np\n", "from albumentations.pytorch import ToTensorV2\n", + "\n", + "from datamint.entities.annotations import ImageSegmentation\n", "from datamint.mlflow.flavors.model import DatamintModel, ModelSettings\n", "from datamint.mlflow.flavors.task_type import TaskType\n", - "from datamint.entities.annotations import ImageSegmentation\n", - "import cv2\n", - "\n", - "_debug('datamint.api')\n", "\n", "\n", "class SegmentationAdapter(DatamintModel):\n", @@ -282,10 +280,12 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint.entities.resource import LocalResource\n", - "from PIL import Image\n", "import io\n", "\n", + "from PIL import Image\n", + "\n", + "from datamint.entities.resource import LocalResource\n", + "\n", "# Create a dummy in-memory PNG as a LocalResource\n", "buf = io.BytesIO()\n", "\n", @@ -346,6 +346,7 @@ "outputs": [], "source": [ "import mlflow\n", + "\n", "import datamint.mlflow as datamint_mlflow\n", "from datamint.mlflow.flavors import log_model\n", "\n", @@ -474,7 +475,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started\")\n", + "print(\"Deployment job started\")\n", "print(f\"Job ID : {job.id}\")\n", "print(f\"Status : {job.status}\")" ] diff --git a/notebooks/05_deployment/03_validate_model.ipynb b/notebooks/05_deployment/03_validate_model.ipynb index 2365e500..1376a597 100644 --- a/notebooks/05_deployment/03_validate_model.ipynb +++ b/notebooks/05_deployment/03_validate_model.ipynb @@ -35,11 +35,9 @@ "metadata": {}, "outputs": [], "source": [ - "import datamint.mlflow.flavors.datamint_flavor as datamint_flavor\n", - "from datamint import validate_model\n", + "from datamint import Api, validate_model\n", "from datamint.dataset import build_dataset\n", - "from datamint import Api\n", - "\n", + "from datamint.mlflow.flavors import datamint_flavor\n", "\n", "PROJECT_NAME = \"bccd_detection\"\n", "MODEL_NAME = PROJECT_NAME\n", @@ -69,7 +67,12 @@ "id": "e5f6a7b8", "metadata": {}, "outputs": [], - "source": "dataset = build_dataset(project=PROJECT_NAME, allow_external_annotations=True)\n\nprint(f\"Resources : {len(dataset.resources)}\")\nprint(f\"Box labels: {dataset.box_labels_set}\")" + "source": [ + "dataset = build_dataset(project=PROJECT_NAME, allow_external_annotations=True)\n", + "\n", + "print(f\"Resources : {len(dataset.resources)}\")\n", + "print(f\"Box labels: {dataset.box_labels_set}\")" + ] }, { "cell_type": "markdown", @@ -131,7 +134,10 @@ "id": "c9d0e1f2", "metadata": {}, "outputs": [], - "source": "report = validate_model(model, dataset=dataset, n_samples=2)\nprint(report)" + "source": [ + "report = validate_model(model, dataset=dataset, n_samples=2)\n", + "print(report)" + ] }, { "cell_type": "markdown", @@ -166,7 +172,12 @@ "id": "b4c5d6e7", "metadata": {}, "outputs": [], - "source": "samples = list(dataset.resources[:3])\n\nreport = validate_model(model, sample_input=samples)\nprint(report)" + "source": [ + "samples = list(dataset.resources[:3])\n", + "\n", + "report = validate_model(model, sample_input=samples)\n", + "print(report)" + ] } ], "metadata": { diff --git a/notebooks/05_deployment/04_predict_images_volumes_and_videos.ipynb b/notebooks/05_deployment/04_predict_images_volumes_and_videos.ipynb index 46c9151b..e6573e63 100644 --- a/notebooks/05_deployment/04_predict_images_volumes_and_videos.ipynb +++ b/notebooks/05_deployment/04_predict_images_volumes_and_videos.ipynb @@ -150,8 +150,8 @@ } ], "source": [ - "from datamint import Api\n", "import datamint.mlflow as datamint_mlflow\n", + "from datamint import Api\n", "\n", "PROJECT_NAME = 'deeplabv3plus_Segmentation_Tutorial'\n", "MODEL_NAME = 'deeplabv3plus_Segmentation_Tutorial'\n", @@ -312,9 +312,12 @@ } ], "source": [ - "import numpy as np\n", + "import os\n", + "import tempfile\n", + "\n", "import nibabel as nib\n", - "import tempfile, os\n", + "import numpy as np\n", + "\n", "from datamint.entities.resource import LocalResource\n", "\n", "# 64 x 64 x 10 volume — 10 axial slices\n", diff --git a/notebooks/06_end_to_end/full_3d/01_synapse_unetrpp.ipynb b/notebooks/06_end_to_end/full_3d/01_synapse_unetrpp.ipynb index ce0850de..56d1fe72 100644 --- a/notebooks/06_end_to_end/full_3d/01_synapse_unetrpp.ipynb +++ b/notebooks/06_end_to_end/full_3d/01_synapse_unetrpp.ipynb @@ -324,8 +324,8 @@ "outputs": [], "source": [ "import h5py\n", - "import numpy as np\n", "import nibabel as nib\n", + "import numpy as np\n", "\n", "NII_DIR = DATA_DIR / 'nifti' / 'images'\n", "LABEL_DIR = DATA_DIR / 'nifti' / 'labels'\n", @@ -532,6 +532,7 @@ "outputs": [], "source": [ "import os\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "from datamint.lightning import UNETRPPTrainer\n", @@ -661,9 +662,9 @@ } ], "source": [ - "import torch\n", - "import numpy as np\n", "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "import torch\n", "\n", "model = trainer.model\n", "model.eval()\n", @@ -815,7 +816,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")" ] diff --git a/notebooks/06_end_to_end/full_3d/02_synapse_nnunet.ipynb b/notebooks/06_end_to_end/full_3d/02_synapse_nnunet.ipynb index 3e5b0ce4..d118ffb7 100644 --- a/notebooks/06_end_to_end/full_3d/02_synapse_nnunet.ipynb +++ b/notebooks/06_end_to_end/full_3d/02_synapse_nnunet.ipynb @@ -337,8 +337,8 @@ ], "source": [ "import h5py\n", - "import numpy as np\n", "import nibabel as nib\n", + "import numpy as np\n", "\n", "NII_DIR = DATA_DIR / 'nifti' / 'images'\n", "LABEL_DIR = DATA_DIR / 'nifti' / 'labels'\n", @@ -568,6 +568,7 @@ "outputs": [], "source": [ "import os\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "from datamint.lightning import NNUNetTrainer\n", @@ -630,6 +631,8 @@ } ], "source": [ + "from pathlib import Path\n", + "\n", "bridge = results['bridge']\n", "model_name = results['model_name']\n", "\n", @@ -637,8 +640,6 @@ "print(f\"Fold output dir : {bridge.output_folder}\")\n", "print(f\"Predictions dir : {bridge.output_folder_base}/predictions_test\")\n", "\n", - "# Check that the final checkpoint exists\n", - "from pathlib import Path\n", "final_ckpt = Path(bridge.output_folder) / 'checkpoint_final.pth'\n", "print(f\"Final checkpoint : {'✓ exists' if final_ckpt.exists() else '✗ not found'}\")\n", "\n", @@ -711,9 +712,10 @@ "source": [ "%matplotlib inline\n", "from pathlib import Path\n", + "\n", + "import matplotlib.pyplot as plt\n", "import nibabel as nib\n", "import numpy as np\n", - "import matplotlib.pyplot as plt\n", "\n", "bridge = results['bridge']\n", "pred_dir = Path(bridge.output_folder_base) / 'predictions_test'\n", @@ -729,7 +731,7 @@ " # Locate the exported test images so we can display CT + prediction side-by-side.\n", " # The exporter writes them to: nnunet_work_dir/raw/Dataset{id}_{name}/imagesTs/\n", " raw_dir = trainer.nnunet_work_dir / 'raw'\n", - " dataset_dir = sorted(raw_dir.glob('Dataset*'))[0]\n", + " dataset_dir = min(raw_dir.glob('Dataset*'))\n", " imagesTs_dir = dataset_dir / 'imagesTs'\n", "\n", " SYNAPSE_CLASSES = {\n", @@ -765,12 +767,12 @@ " pred_slice = pred_vol[:, :, s]\n", "\n", " axes[row, 0].imshow(ct_slice, cmap='gray')\n", - " axes[row, 0].set_title(f'CT — axial')\n", + " axes[row, 0].set_title('CT — axial')\n", " axes[row, 0].axis('off')\n", "\n", " axes[row, 1].imshow(ct_slice, cmap='gray')\n", " axes[row, 1].imshow(pred_slice, cmap=CMAP, alpha=0.5, vmin=0, vmax=num_classes)\n", - " axes[row, 1].set_title(f'nnU-Net prediction')\n", + " axes[row, 1].set_title('nnU-Net prediction')\n", " axes[row, 1].axis('off')\n", "\n", " handles = [plt.Rectangle((0, 0), 1, 1, color=CMAP(i + 1)) for i in range(num_classes)]\n", @@ -812,7 +814,7 @@ " with_gpu=True,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")" ] @@ -901,7 +903,53 @@ "id": "315e1ad6", "metadata": {}, "outputs": [], - "source": "%matplotlib inline\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Fetch CT volume\nct_nifti = r.fetch_file_data(use_cache=True, auto_convert=True)\nct_vol = ct_nifti.get_fdata()\n\n# Fetch the predicted segmentation (most recent annotation on this resource)\nannotations = list(r.fetch_annotations())\npred_ann = annotations[-1]\npred_vol = pred_ann.fetch_file_data(auto_convert=True, use_cache=True)\n\nSYNAPSE_CLASSES = {\n 1: 'aorta', 2: 'gallbladder', 3: 'spleen',\n 4: 'left_kidney', 5: 'right_kidney', 6: 'liver',\n 7: 'stomach', 8: 'pancreas',\n}\nnum_classes = len(SYNAPSE_CLASSES)\nCMAP = plt.get_cmap('tab10', num_classes + 1)\n\n# nnU-Net volumes are (H, W, D) — pick 3 axial slices\nD = ct_vol.shape[2]\nslice_indices = [D // 4, D // 2, 3 * D // 4]\n\nfig, axes = plt.subplots(len(slice_indices), 2, figsize=(10, 4 * len(slice_indices)))\nfig.suptitle(f'nnU-Net inference result — {r.filename}', fontsize=13)\n\nfor row, s in enumerate(slice_indices):\n ct_slice = ct_vol[:, :, s]\n pred_slice = pred_vol[:, :, s]\n\n axes[row, 0].imshow(ct_slice, cmap='gray')\n axes[row, 0].set_title(f'CT — axial slice {s}')\n axes[row, 0].axis('off')\n\n axes[row, 1].imshow(ct_slice, cmap='gray')\n axes[row, 1].imshow(pred_slice, cmap=CMAP, alpha=0.5, vmin=0, vmax=num_classes)\n axes[row, 1].set_title('nnU-Net Prediction')\n axes[row, 1].axis('off')\n\nhandles = [plt.Rectangle((0, 0), 1, 1, color=CMAP(i + 1)) for i in range(num_classes)]\nfig.legend(handles, list(SYNAPSE_CLASSES.values()), loc='lower center', ncol=4, fontsize=9, title='Organ classes')\nplt.tight_layout()\nplt.show()" + "source": [ + "%matplotlib inline\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", + "\n", + "# Fetch CT volume\n", + "ct_nifti = r.fetch_file_data(use_cache=True, auto_convert=True)\n", + "ct_vol = ct_nifti.get_fdata()\n", + "\n", + "# Fetch the predicted segmentation (most recent annotation on this resource)\n", + "annotations = list(r.fetch_annotations())\n", + "pred_ann = annotations[-1]\n", + "pred_vol = pred_ann.fetch_file_data(auto_convert=True, use_cache=True)\n", + "\n", + "SYNAPSE_CLASSES = {\n", + " 1: 'aorta', 2: 'gallbladder', 3: 'spleen',\n", + " 4: 'left_kidney', 5: 'right_kidney', 6: 'liver',\n", + " 7: 'stomach', 8: 'pancreas',\n", + "}\n", + "num_classes = len(SYNAPSE_CLASSES)\n", + "CMAP = plt.get_cmap('tab10', num_classes + 1)\n", + "\n", + "# nnU-Net volumes are (H, W, D) — pick 3 axial slices\n", + "D = ct_vol.shape[2]\n", + "slice_indices = [D // 4, D // 2, 3 * D // 4]\n", + "\n", + "fig, axes = plt.subplots(len(slice_indices), 2, figsize=(10, 4 * len(slice_indices)))\n", + "fig.suptitle(f'nnU-Net inference result — {r.filename}', fontsize=13)\n", + "\n", + "for row, s in enumerate(slice_indices):\n", + " ct_slice = ct_vol[:, :, s]\n", + " pred_slice = pred_vol[:, :, s]\n", + "\n", + " axes[row, 0].imshow(ct_slice, cmap='gray')\n", + " axes[row, 0].set_title(f'CT — axial slice {s}')\n", + " axes[row, 0].axis('off')\n", + "\n", + " axes[row, 1].imshow(ct_slice, cmap='gray')\n", + " axes[row, 1].imshow(pred_slice, cmap=CMAP, alpha=0.5, vmin=0, vmax=num_classes)\n", + " axes[row, 1].set_title('nnU-Net Prediction')\n", + " axes[row, 1].axis('off')\n", + "\n", + "handles = [plt.Rectangle((0, 0), 1, 1, color=CMAP(i + 1)) for i in range(num_classes)]\n", + "fig.legend(handles, list(SYNAPSE_CLASSES.values()), loc='lower center', ncol=4, fontsize=9, title='Organ classes')\n", + "plt.tight_layout()\n", + "plt.show()" + ] } ], "metadata": { @@ -925,4 +973,4 @@ }, "nbformat": 4, "nbformat_minor": 5 -} \ No newline at end of file +} diff --git a/notebooks/06_end_to_end/slice_based/01_fracatlas_classification.ipynb b/notebooks/06_end_to_end/slice_based/01_fracatlas_classification.ipynb index 418f9605..ed2fa79c 100644 --- a/notebooks/06_end_to_end/slice_based/01_fracatlas_classification.ipynb +++ b/notebooks/06_end_to_end/slice_based/01_fracatlas_classification.ipynb @@ -248,9 +248,10 @@ "metadata": {}, "outputs": [], "source": [ - "import requests\n", - "import zipfile\n", "import os\n", + "import zipfile\n", + "\n", + "import requests\n", "\n", "# Retrieve and download FracAtlas dataset from Figshare\n", "# It might take a while depending on your internet connection. ~50 seconds on a 100Mbps connection\n", @@ -305,8 +306,8 @@ "metadata": {}, "outputs": [], "source": [ - "from pathlib import Path\n", "import os\n", + "from pathlib import Path\n", "\n", "# get all non-fractured images\n", "non_fractured_root_path = Path('FracAtlas/FracAtlas/images/Non_fractured/')\n", @@ -477,8 +478,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint.lightning import EfficientNetV2Trainer\n", "import os\n", + "\n", + "from datamint.lightning import EfficientNetV2Trainer\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "trainer = EfficientNetV2Trainer(\n", @@ -612,10 +615,11 @@ "metadata": {}, "outputs": [], "source": [ + "import torch\n", + "from torchvision.models import resnet18\n", + "\n", "from datamint.lightning import ImageClassificationTrainer\n", "from datamint.lightning.trainers.lightning_modules import ClassificationModule\n", - "from torchvision.models import resnet18\n", - "import torch\n", "\n", "\n", "class FracAtlasClassifier(ClassificationModule):\n", @@ -747,7 +751,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")" ] diff --git a/notebooks/06_end_to_end/slice_based/02_busi_segmentation.ipynb b/notebooks/06_end_to_end/slice_based/02_busi_segmentation.ipynb index ebb3d1cb..8c6264ae 100644 --- a/notebooks/06_end_to_end/slice_based/02_busi_segmentation.ipynb +++ b/notebooks/06_end_to_end/slice_based/02_busi_segmentation.ipynb @@ -224,10 +224,11 @@ "outputs": [], "source": [ "import os\n", - "import requests\n", "import zipfile\n", "from pathlib import Path\n", "\n", + "import requests\n", + "\n", "BUSI_URL = \"https://www.kaggle.com/api/v1/datasets/download/sabahesaraki/breast-ultrasound-images-dataset\"\n", "DATA_DIR = Path(\"/tmp/BUSI_dataset\")\n", "\n", @@ -238,8 +239,7 @@ " zip_path = DATA_DIR / \"Dataset_BUSI.zip\"\n", " DATA_DIR.mkdir(parents=True, exist_ok=True)\n", " with open(zip_path, 'wb') as f:\n", - " for chunk in response.iter_content(chunk_size=8192):\n", - " f.write(chunk)\n", + " f.writelines(response.iter_content(chunk_size=8192))\n", " print(\"Extracting...\")\n", " with zipfile.ZipFile(zip_path, 'r') as zip_ref:\n", " zip_ref.extractall(DATA_DIR)\n", @@ -436,8 +436,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint.lightning import DeepLabV3PlusTrainer\n", "import os\n", + "\n", + "from datamint.lightning import DeepLabV3PlusTrainer\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "trainer = DeepLabV3PlusTrainer(\n", @@ -465,8 +467,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint.lightning import UNetPPTrainer\n", "import os\n", + "\n", + "from datamint.lightning import UNetPPTrainer\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "trainer = UNetPPTrainer(\n", @@ -494,8 +498,10 @@ "metadata": {}, "outputs": [], "source": [ - "from datamint.lightning import TransUNetTrainer\n", "import os\n", + "\n", + "from datamint.lightning import TransUNetTrainer\n", + "\n", "os.environ['MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR'] = 'false'\n", "\n", "trainer = TransUNetTrainer(\n", @@ -584,9 +590,10 @@ "metadata": {}, "outputs": [], "source": [ + "import segmentation_models_pytorch as smp\n", + "\n", "from datamint.lightning import SemanticSegmentation2DTrainer\n", "from datamint.lightning.trainers.lightning_modules import SegmentationModule\n", - "import segmentation_models_pytorch as smp\n", "\n", "\n", "class MyCustomSegModel(SegmentationModule):\n", @@ -790,12 +797,13 @@ } ], "source": [ - "import torch\n", "import numpy as np\n", + "import torch\n", "from matplotlib import pyplot as plt\n", - "from datamint.utils.visualization import show, draw_masks\n", "from torchmetrics.functional.segmentation import mean_iou\n", "\n", + "from datamint.utils.visualization import draw_masks, show\n", + "\n", "model = trainer.model\n", "model.eval()\n", "\n", @@ -1012,7 +1020,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")" ] diff --git a/notebooks/06_end_to_end/slice_based/03_bccd_detection.ipynb b/notebooks/06_end_to_end/slice_based/03_bccd_detection.ipynb index a3c7f07a..258869b9 100644 --- a/notebooks/06_end_to_end/slice_based/03_bccd_detection.ipynb +++ b/notebooks/06_end_to_end/slice_based/03_bccd_detection.ipynb @@ -282,7 +282,10 @@ "@dataclass\n", "class Box:\n", " label: str\n", - " x1: float; y1: float; x2: float; y2: float\n", + " x1: float\n", + " y1: float\n", + " x2: float\n", + " y2: float\n", "\n", "\n", "@dataclass\n", @@ -340,9 +343,9 @@ } ], "source": [ - "import numpy as np\n", - "import matplotlib.pyplot as plt\n", "import matplotlib.patches as mpatches\n", + "import matplotlib.pyplot as plt\n", + "import numpy as np\n", "from PIL import Image\n", "\n", "COLORS = {\"RBC\": \"#e74c3c\", \"WBC\": \"#3498db\", \"Platelets\": \"#2ecc71\"}\n", @@ -504,9 +507,10 @@ "metadata": {}, "outputs": [], "source": [ + "import os\n", + "\n", "from datamint.lightning import YOLOXTrainer\n", "\n", - "import os\n", "os.environ[\"MLFLOW_ENABLE_ARTIFACTS_PROGRESS_BAR\"] = \"false\"\n", "\n", "trainer = YOLOXTrainer(\n", @@ -611,11 +615,9 @@ } ], "source": [ + "import matplotlib.pyplot as plt\n", "import torch\n", "from yolox.utils import postprocess\n", - "import matplotlib.pyplot as plt\n", - "import matplotlib.patches as mpatches\n", - "import numpy as np\n", "\n", "COLORS = {\"RBC\": \"#e74c3c\", \"WBC\": \"#3498db\", \"Platelets\": \"#2ecc71\"}\n", "\n", @@ -705,7 +707,7 @@ " with_gpu=False,\n", ")\n", "\n", - "print(f\"Deployment job started!\")\n", + "print(\"Deployment job started!\")\n", "print(f\"Job ID: {job.id}\")\n", "print(f\"Status: {job.status}\")" ] diff --git a/pyproject.toml b/pyproject.toml index aa31368f..6c854a87 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -68,6 +68,7 @@ nnunetv2 = { version = ">=2.4,<3.0", optional = true } filelock = { version = ">=3.0", optional = true } yolox-datamint = { version = ">=0.3.1", optional = true } h5py = { version = ">=3.0", optional = true } +ruff = { version = "^0.16.1", optional = true } [tool.poetry.group.dev.dependencies] # for `poetry install` @@ -76,11 +77,12 @@ pytest-cov = "^7.1.0" responses = "^0.20.0" aioresponses = "^0.7.0" respx = ">=0.22.0" +ruff = "^0.16.1" # Extra dependencies for docs [tool.poetry.extras] docs = ["sphinx", "sphinx_rtd_theme", "sphinx-tabs", "sphinx-rtd-dark-mode", "setuptools"] -dev = ["pytest", "pytest-cov", "responses", "aioresponses", "respx"] +dev = ["pytest", "pytest-cov", "responses", "aioresponses", "respx", "ruff"] nnunet = ["nnunetv2", "filelock"] detection = ["yolox-datamint"] examples = ["h5py"] @@ -115,3 +117,14 @@ https = "datamint.mlflow.models.datamint_model_store:DatamintModelRegistryStore" [tool.pytest.ini_options] testpaths = ["tests"] addopts = "-v --tb=short" + +[tool.ruff] +target-version = "py310" +line-length = 120 + +[tool.ruff.lint] +# Narrow on purpose: only categories already mostly clean in this codebase. +# BLE (blind except), S (bandit), TRY (tryceratops), SIM (simplify) are +# left out +select = ["E4", "E7", "E9", "F", "RUF"] +ignore = ["RUF001", "RUF002", "RUF003"] diff --git a/tests/conftest.py b/tests/conftest.py index 31fde300..1963b1ec 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,7 +8,6 @@ from datamint.api.base_api import ApiConfig - TEST_URL = "https://test-url.com" diff --git a/tests/test_annotation_agreement.py b/tests/test_annotation_agreement.py index 10d53cd8..4d84b164 100644 --- a/tests/test_annotation_agreement.py +++ b/tests/test_annotation_agreement.py @@ -1,7 +1,11 @@ import numpy as np import pytest -from datamint.entities.annotations import BoxAnnotation, ImageClassification, ImageSegmentation +from datamint.entities.annotations import ( + BoxAnnotation, + ImageClassification, + ImageSegmentation, +) from datamint.utils.annotation_agreement import ( cohen_kappa, compute_agreement, diff --git a/tests/test_api_handler.py b/tests/test_api_handler.py index 592f20a8..8dee9899 100644 --- a/tests/test_api_handler.py +++ b/tests/test_api_handler.py @@ -1,14 +1,16 @@ -import pytest +from typing import IO from unittest.mock import patch -from datamint.api.client import Api -from datamint.exceptions import DatamintException -import respx + +import httpx +import numpy as np import pydicom -from pydicom.data import get_testdata_files +import pytest +import respx from aiohttp import FormData -from typing import IO -import numpy as np -import httpx +from pydicom.data import get_testdata_files + +from datamint.api.client import Api +from datamint.exceptions import DatamintException # pytest tests --log-cli-level=INFO @@ -138,11 +140,11 @@ def sample_2dmask1(self) -> np.ndarray: @respx.mock @patch('os.getenv') def test_api_handler_init(self, mock_getenv, get_projects_sample: dict): - api_handler = Api(check_connection=False) + Api(check_connection=False) ### Test wrong url ### with pytest.raises(DatamintException): - api_handler = Api(server_url='wrong', check_connection=True) + Api(server_url='wrong', check_connection=True) @respx.mock def test_check_connection_is_cached_and_resets_on_config_change(self, get_projects_sample: dict): diff --git a/tests/test_auto_config_mlflow.py b/tests/test_auto_config_mlflow.py index 61ae4f10..e6910e94 100644 --- a/tests/test_auto_config_mlflow.py +++ b/tests/test_auto_config_mlflow.py @@ -2,12 +2,12 @@ Test that MLflow environment is properly configured regardless of import order. """ -import pytest -from unittest.mock import patch -import sys import os +import sys from itertools import permutations +from unittest.mock import patch +import pytest # Test configuration values TEST_API_URL = "http://localhost:3001" @@ -113,7 +113,7 @@ def test_mlflow_uri_patching_all_import_orders(import_order, mock_datamint_confi # Execute imports in the specified order namespace = {} for idx in import_order: - name, import_stmt = IMPORT_STATEMENTS[idx] + _name, import_stmt = IMPORT_STATEMENTS[idx] exec(import_stmt, namespace) # Get MLFlowLogger from namespace (it should have been imported) diff --git a/tests/test_datamint_config.py b/tests/test_datamint_config.py index 7bd11bd7..1c37315b 100644 --- a/tests/test_datamint_config.py +++ b/tests/test_datamint_config.py @@ -88,7 +88,6 @@ def test_config_imports_no_torch(self) -> None: initial_modules = set(sys.modules.keys()) # Import the config module - from datamint.client_cmd_tools.datamint_config import main # Check new modules that were imported new_modules = set(sys.modules.keys()) - initial_modules @@ -128,7 +127,6 @@ def test_config_tool_startup_time(self) -> None: import time start_time = time.time() - from datamint.client_cmd_tools.datamint_config import main end_time = time.time() startup_time = end_time - start_time @@ -149,7 +147,6 @@ def test_config_persistence(self, mock_set, mock_get) -> None: mock_get.return_value = None # Test setting and getting a value - test_key = 'test_config_key' test_value = 'test_config_value' with patch('sys.argv', ['datamint-config', '--api-key', test_value]): @@ -161,7 +158,7 @@ def test_config_persistence(self, mock_set, mock_get) -> None: def test_environment_variable_integration(self) -> None: """Test config integration with environment variables.""" import os - from datamint import configs + # Test environment variable fallback test_key = 'DATAMINT_TEST_VAR' diff --git a/tests/test_dataset_patient_split.py b/tests/test_dataset_patient_split.py index 3d7a9542..b2f1c48b 100644 --- a/tests/test_dataset_patient_split.py +++ b/tests/test_dataset_patient_split.py @@ -3,7 +3,6 @@ from datamint.dataset.base import DatamintBaseDataset - # --------------------------------------------------------------------------- # Test setup # --------------------------------------------------------------------------- diff --git a/tests/test_detection_trainer.py b/tests/test_detection_trainer.py index b5919fc0..d4d306c9 100644 --- a/tests/test_detection_trainer.py +++ b/tests/test_detection_trainer.py @@ -1,13 +1,12 @@ """Tests for DetectionTrainer abstract base.""" -import pytest from unittest.mock import MagicMock, patch import albumentations as A from albumentations.pytorch import ToTensorV2 -from datamint.lightning.trainers.detection_trainer import DetectionTrainer from datamint.dataset.image_dataset import ImageDataset, detection_collate_fn from datamint.lightning.datamodule import DatamintDataModule +from datamint.lightning.trainers.detection_trainer import DetectionTrainer class _ConcreteDetectionTrainer(DetectionTrainer): diff --git a/tests/test_imports.py b/tests/test_imports.py index 174c0007..acbb0af1 100644 --- a/tests/test_imports.py +++ b/tests/test_imports.py @@ -2,9 +2,10 @@ Test module to verify that all important datamint modules can be imported successfully. This helps catch import issues early and ensures the package structure is correct. """ -import pytest import logging +import pytest + # Set up logging to capture any import warnings logging.basicConfig(level=logging.WARNING) _LOGGER = logging.getLogger(__name__) diff --git a/tests/test_nnunet_data_export.py b/tests/test_nnunet_data_export.py index 0842038c..499cf370 100644 --- a/tests/test_nnunet_data_export.py +++ b/tests/test_nnunet_data_export.py @@ -2,13 +2,18 @@ Testing the DatamintToNNUNetExporter class to ensure it correctly writes dataset.json, preserves voxel spacing, merges segmentations, and creates the expected directory structure for nnUNet. """ import pytest + pytest.importorskip("nnunetv2", minversion="2.4") import json -import numpy as np +from unittest.mock import MagicMock + import nibabel as nib +import numpy as np import pytest -from unittest.mock import MagicMock -from datamint.lightning.trainers.specialized.nnunet.data_export import DatamintToNNUNetExporter + +from datamint.lightning.trainers.specialized.nnunet.data_export import ( + DatamintToNNUNetExporter, +) def _make_nifti(shape=(64, 64, 32), zooms=(1.5, 1.5, 2.0)) -> nib.Nifti1Image: @@ -63,11 +68,15 @@ def test_export_image_preserves_voxel_spacing(tmp_path): def test_merge_segmentations_highest_class_wins(tmp_path): shape = (32, 32, 16) - liver = np.zeros(shape, dtype=np.int32); liver[5:15, 5:15, 2:8] = 1 - tumor = np.zeros(shape, dtype=np.int32); tumor[10:20, 10:20, 4:10] = 2 + liver = np.zeros(shape, dtype=np.int32) + liver[5:15, 5:15, 2:8] = 1 + tumor = np.zeros(shape, dtype=np.int32) + tumor[10:20, 10:20, 4:10] = 2 - liver_seg = MagicMock(); liver_seg.fetch_file_data.return_value = liver - tumor_seg = MagicMock(); tumor_seg.fetch_file_data.return_value = tumor + liver_seg = MagicMock() + liver_seg.fetch_file_data.return_value = liver + tumor_seg = MagicMock() + tumor_seg.fetch_file_data.return_value = tumor exp = DatamintToNNUNetExporter(tmp_path, dataset_id=1, dataset_name='CTLiver') merged = exp._merge_segmentations([liver_seg, tumor_seg]) @@ -78,8 +87,10 @@ def test_merge_segmentations_highest_class_wins(tmp_path): def test_merge_segmentations_warns_on_overlap(tmp_path): shape = (32, 32, 16) - seg1 = MagicMock(); seg1.fetch_file_data.return_value = np.ones(shape, dtype=np.int32) - seg2 = MagicMock(); seg2.fetch_file_data.return_value = np.ones(shape, dtype=np.int32) * 2 + seg1 = MagicMock() + seg1.fetch_file_data.return_value = np.ones(shape, dtype=np.int32) + seg2 = MagicMock() + seg2.fetch_file_data.return_value = np.ones(shape, dtype=np.int32) * 2 exp = DatamintToNNUNetExporter(tmp_path, dataset_id=1, dataset_name='CTLiver') with pytest.warns(UserWarning, match='overlap'): diff --git a/tests/test_nnunet_data_import.py b/tests/test_nnunet_data_import.py index 8077f7b9..dc4cc4a5 100644 --- a/tests/test_nnunet_data_import.py +++ b/tests/test_nnunet_data_import.py @@ -2,14 +2,19 @@ Test the NNUNetToDatamintImporter class for importing nnUNet predictions into Datamint. """ import pytest + pytest.importorskip("nnunetv2", minversion="2.4") import json -import numpy as np -import nibabel as nib -import pytest from pathlib import Path from unittest.mock import MagicMock -from datamint.lightning.trainers.specialized.nnunet.data_import import NNUNetToDatamintImporter + +import nibabel as nib +import numpy as np +import pytest + +from datamint.lightning.trainers.specialized.nnunet.data_import import ( + NNUNetToDatamintImporter, +) def _write_pred(path: Path, shape=(64, 64, 32)): diff --git a/tests/test_nnunet_inference_model.py b/tests/test_nnunet_inference_model.py index 396ba79e..bb7a856d 100644 --- a/tests/test_nnunet_inference_model.py +++ b/tests/test_nnunet_inference_model.py @@ -3,14 +3,20 @@ We also check that temporary directories are cleaned up after prediction. """ import pytest + pytest.importorskip("nnunetv2", minversion="2.4") import sys -import numpy as np -import nibabel as nib from pathlib import Path from unittest.mock import MagicMock, patch -from datamint.lightning.trainers.specialized.nnunet.inference_model import NNUNetInferenceModel + +import nibabel as nib +import numpy as np + +from datamint.lightning.trainers.specialized.nnunet.inference_model import ( + NNUNetInferenceModel, +) + # Fake nnunetv2 submodules so patch() can resolve them def _mock_nnunetv2(): diff --git a/tests/test_nnunet_integration.py b/tests/test_nnunet_integration.py index 3bb622eb..8c5babb6 100644 --- a/tests/test_nnunet_integration.py +++ b/tests/test_nnunet_integration.py @@ -10,7 +10,9 @@ - Enough disk space for nnUNet preprocessing (~2–5 GB for a small dataset) """ import random + import pytest + pytest.importorskip("nnunetv2", minversion="2.4") from pathlib import Path @@ -20,9 +22,10 @@ @pytest.mark.integration def test_full_nnunet_pipeline(): + import mlflow + from datamint import Api from datamint.lightning import NNUNetTrainer - import mlflow api = Api() diff --git a/tests/test_nnunet_trainer.py b/tests/test_nnunet_trainer.py index 1fcc1c48..4e4248b1 100644 --- a/tests/test_nnunet_trainer.py +++ b/tests/test_nnunet_trainer.py @@ -3,12 +3,16 @@ called in the correct order during fit(). We also test the dataset ID assignment logic and that the expected files are written during fingerprinting/planning. """ import os +import re import sys -import pytest from unittest.mock import MagicMock, patch + +import pytest + import datamint.lightning.trainers.specialized.nnunet.trainer as _trainer_mod from datamint.lightning.trainers.specialized.nnunet.trainer import NNUNetTrainer + # Fake nnunetv2 submodules so patch() can resolve them when the trainer imports them. def _mock_nnunetv2(): root = MagicMock() @@ -111,7 +115,7 @@ def test_fingerprint_raises_if_json_not_written(trainer): '.fingerprint_extractor.DatasetFingerprintExtractor') as MockFP, \ patch('nnunetv2.experiment_planning.experiment_planners.default_experiment_planner.ExperimentPlanner'): MockFP.return_value.run.return_value = None - with pytest.raises(RuntimeError, match='dataset_fingerprint.json'): + with pytest.raises(RuntimeError, match=re.escape('dataset_fingerprint.json')): trainer._run_fingerprint_and_plan(dataset_id=1) diff --git a/tests/test_nnunet_trainer_bridge.py b/tests/test_nnunet_trainer_bridge.py deleted file mode 100644 index c3b48ae9..00000000 --- a/tests/test_nnunet_trainer_bridge.py +++ /dev/null @@ -1,128 +0,0 @@ -""" Test the _DatamintNNUNetTrainer bridge class that extends nnUNetTrainer to log checkpoints and validation summaries to MLflow. We verify that checkpoints are logged as artifacts, that validation metrics are -logged correctly, and that the version guard prevents usage with old nnunetv2 versions.""" -import json -import importlib.metadata as _meta -import pytest - -try: - _ver = tuple(int(x) for x in _meta.version('nnunetv2').split('.')[:2]) - if not ((2, 4) <= _ver < (3, 0)): - pytest.skip("nnunetv2>=2.4,<3.0 required", allow_module_level=True) -except _meta.PackageNotFoundError: - pytest.skip("nnunetv2 not installed", allow_module_level=True) - -from unittest.mock import MagicMock, patch -import sys - -# Mock nnunetv2 BEFORE importing the bridge so tests run without nnunetv2 installed. -class _FakeNNUNetTrainer: - def save_checkpoint(self, *a, **kw): pass - def perform_actual_validation(self, *a, **kw): pass - def print_to_log_file(self, *a, **kw): pass - -_nnunetv2_mock = MagicMock() -_nnunetv2_mock.__version__ = '2.4.0' -sys.modules.setdefault('nnunetv2', _nnunetv2_mock) -sys.modules.setdefault('nnunetv2.training', MagicMock()) -sys.modules.setdefault('nnunetv2.training.nnUNetTrainer', MagicMock()) - -_trainer_module_mock = MagicMock() -_trainer_module_mock.nnUNetTrainer = _FakeNNUNetTrainer -sys.modules.setdefault('nnunetv2.training.nnUNetTrainer.nnUNetTrainer', _trainer_module_mock) - -from datamint.lightning.trainers.specialized.nnunet._nnunet_trainer_bridge import ( - _DatamintNNUNetTrainer, - _MLflowLogger, - _SKIP_METRIC_KEYS, -) - - -@pytest.fixture() -def bridge(): - b = _DatamintNNUNetTrainer.__new__(_DatamintNNUNetTrainer) - b._best_checkpoint_path = None - b.current_epoch = 5 - return b - - -@pytest.fixture() -def mlflow_logger(): - return _MLflowLogger() - - -# -- Tests for _MLflowLogger -def test_mlflow_logger_log_emits_metric(mlflow_logger): - with patch('mlflow.log_metric') as mock_log: - mlflow_logger.log('train_losses', 0.42, step=3) - mock_log.assert_called_once_with('train_losses', 0.42, step=3) - - -def test_mlflow_logger_skips_timestamp_keys(mlflow_logger): - with patch('mlflow.log_metric') as mock_log: - for key in _SKIP_METRIC_KEYS: - mlflow_logger.log(key, 12345.0, step=0) - mock_log.assert_not_called() - - -def test_mlflow_logger_log_summary(mlflow_logger): - with patch('mlflow.log_metric') as mock_log: - mlflow_logger.log_summary('final_val/foreground_dice', 0.88) - mock_log.assert_called_once_with('final_val/foreground_dice', 0.88) - - -def test_mlflow_logger_update_config_logs_params(mlflow_logger): - with patch('mlflow.log_params') as mock_params: - mlflow_logger.update_config({'lr': 0.01, 'epochs': 100, 'name': 'test'}) - mock_params.assert_called_once_with({'lr': 0.01, 'epochs': 100, 'name': 'test'}) - - -def test_mlflow_logger_skips_non_scalar_config(mlflow_logger): - with patch('mlflow.log_params') as mock_params: - mlflow_logger.update_config({'lr': 0.01, 'bad': [1, 2, 3]}) - mock_params.assert_called_once_with({'lr': 0.01}) - - -# -- Tests for _DatamintNNUNetTrainer -def test_save_checkpoint_records_path(bridge, tmp_path): - ckpt = tmp_path / 'checkpoint_best.pth'; ckpt.touch() - with patch('mlflow.log_artifact'), \ - patch.object(_DatamintNNUNetTrainer, '_super_save_checkpoint'): - bridge.save_checkpoint(str(ckpt)) - assert bridge._best_checkpoint_path == ckpt - - -def test_save_checkpoint_logs_artifact(bridge, tmp_path): - ckpt = tmp_path / 'checkpoint_best.pth'; ckpt.touch() - with patch('mlflow.log_artifact') as mock_artifact, \ - patch.object(_DatamintNNUNetTrainer, '_super_save_checkpoint'): - bridge.save_checkpoint(str(ckpt)) - mock_artifact.assert_called_once_with(str(ckpt), artifact_path='nnunet_checkpoints') - - -def test_log_validation_summary_logs_per_class_dice(bridge, tmp_path): - summary = { - 'foreground_mean': 0.85, - 'mean': {'liver': {'Dice': 0.91}, 'tumor': {'Dice': 0.74}}, - } - val_dir = tmp_path / 'validation' - val_dir.mkdir(parents=True) - (val_dir / 'summary.json').write_text(json.dumps(summary)) - bridge.fold = 0 - bridge.output_folder = str(tmp_path) - - with patch('mlflow.log_metric') as mock_log: - bridge._log_validation_summary() - - calls = {c.args[0]: c.args[1] for c in mock_log.call_args_list} - assert calls['val/dice_liver'] == pytest.approx(0.91) - assert calls['val/dice_tumor'] == pytest.approx(0.74) - assert calls['val/dice_mean'] == pytest.approx(0.85) - - -def test_version_guard_raises_on_old_nnunet(): - import importlib - import importlib.metadata as _meta - import datamint.lightning.trainers.specialized.nnunet._nnunet_trainer_bridge as m - with patch.object(_meta, 'version', return_value='1.0.0'): - with pytest.raises(ImportError, match='nnunetv2>=2.4'): - importlib.reload(m) diff --git a/tests/test_upload_validation.py b/tests/test_upload_validation.py index 35063397..5d88507f 100644 --- a/tests/test_upload_validation.py +++ b/tests/test_upload_validation.py @@ -1,13 +1,10 @@ """ Test module for upload functionality validation. """ -import pytest -import tempfile +import logging import os -import json +import tempfile from pathlib import Path -from unittest.mock import patch, MagicMock -import logging _LOGGER = logging.getLogger(__name__) diff --git a/tests/test_yolox_module.py b/tests/test_yolox_module.py deleted file mode 100644 index 28d9474d..00000000 --- a/tests/test_yolox_module.py +++ /dev/null @@ -1,178 +0,0 @@ -"""Tests for YOLOXModule.""" -import sys -import torch -import pytest -from unittest.mock import MagicMock, patch - -# Mock yolox before any import touches it. sys.modules['yolox.models'] must be -# the same object as getattr(yolox_mock, 'models') so that both `import yolox.models` -# (production code) and patch('yolox.models.yolox_s') (tests) resolve to the same object. -_yolox_mock = MagicMock() -sys.modules.setdefault('yolox', _yolox_mock) -sys.modules.setdefault('yolox.models', _yolox_mock.models) -sys.modules.setdefault('yolox.utils', _yolox_mock.utils) - -from datamint.lightning.trainers.lightning_modules.detection_modules.yolox_module import YOLOXModule -from datamint.entities.annotations import BoxAnnotation - -@pytest.fixture() -def module(): - with patch('yolox.models.yolox_s') as MockCtor: - MockCtor.return_value = MagicMock() - m = YOLOXModule(num_classes=2, model_size='s') - m.model = MagicMock() - m.map_metric = None - return m - -# -- testing __init__ and model_size handling -- -def test_invalid_model_size_raises(): - with pytest.raises(ValueError, match='model_size'): - with patch('yolox.models.yolox_bad', create=True): - YOLOXModule(num_classes=1, model_size='bad') - - -def test_model_size_routes_to_correct_variant(): - for size in ('nano', 'tiny', 's', 'm', 'l', 'x'): - target = f'yolox.models.yolox_{size}' - with patch(target) as MockVariant: - MockVariant.return_value = MagicMock() - m = YOLOXModule(num_classes=1, model_size=size) - MockVariant.assert_called_once_with(num_classes=1) - - -# -- testing _build_targets -- - -def test_build_targets_shape(): - boxes = [torch.tensor([[10., 20., 50., 60.]]), torch.zeros(0, 4)] - labels = [torch.tensor([1]), torch.zeros(0, dtype=torch.int64)] - t = YOLOXModule._build_targets(boxes, labels, torch.device('cpu')) - assert t.shape == (2, 1, 5) - - -def test_build_targets_center_form(): - # x1=0, y1=0, x2=10, y2=20 → cx=5, cy=10, w=10, h=20 - boxes = [torch.tensor([[0., 0., 10., 20.]])] - labels = [torch.tensor([0])] - t = YOLOXModule._build_targets(boxes, labels, torch.device('cpu')) - assert t[0, 0, 1].item() == pytest.approx(5.0) # cx - assert t[0, 0, 2].item() == pytest.approx(10.0) # cy - assert t[0, 0, 3].item() == pytest.approx(10.0) # w - assert t[0, 0, 4].item() == pytest.approx(20.0) # h - - -def test_build_targets_class_index(): - boxes = [torch.tensor([[0., 0., 10., 10.]])] - labels = [torch.tensor([3])] - t = YOLOXModule._build_targets(boxes, labels, torch.device('cpu')) - assert t[0, 0, 0].item() == pytest.approx(3.0) - - -def test_build_targets_empty_batch_no_crash(): - boxes = [torch.zeros(0, 4), torch.zeros(0, 4)] - labels = [torch.zeros(0, dtype=torch.int64), torch.zeros(0, dtype=torch.int64)] - t = YOLOXModule._build_targets(boxes, labels, torch.device('cpu')) - - assert t.shape == (2, 1, 5) - assert t.sum().item() == pytest.approx(0.0) - - -# -- testing training_step -- - -def test_training_step_logs_train_loss(module): - module.model.return_value = { - 'total_loss': torch.tensor(1.5), - 'iou_loss': torch.tensor(0.5), - 'cls_loss': torch.tensor(0.3), - 'conf_loss': torch.tensor(0.4), - } - batch = { - 'image': torch.zeros(2, 3, 416, 416), - 'boxes': [torch.zeros(2, 4), torch.zeros(1, 4)], - 'box_labels': [torch.zeros(2, dtype=torch.int64), torch.zeros(1, dtype=torch.int64)], - } - with patch.object(module, 'log') as mock_log: - loss = module.training_step(batch, 0) - - logged_keys = [c.args[0] for c in mock_log.call_args_list] - assert 'train/loss' in logged_keys - assert 'train/iou_loss' in logged_keys - assert 'train/cls_loss' in logged_keys - assert 'train/obj_loss' in logged_keys - assert loss.item() == pytest.approx(1.5) - - -def test_training_step_returns_total_loss(module): - module.model.return_value = { - 'total_loss': torch.tensor(2.0), - 'iou_loss': torch.tensor(0.5), - 'cls_loss': torch.tensor(0.5), - 'conf_loss': torch.tensor(0.5), - } - batch = { - 'image': torch.zeros(1, 3, 416, 416), - 'boxes': [torch.zeros(0, 4)], - 'box_labels': [torch.zeros(0, dtype=torch.int64)], - } - with patch.object(module, 'log'): - loss = module.training_step(batch, 0) - assert loss.item() == pytest.approx(2.0) - - -# -- testing predict_image -- - -def _make_fake_resource(img_np=None): - """Return a mock resource whose fetch_file_data yields a small numpy image.""" - import numpy as np - from unittest.mock import MagicMock - if img_np is None: - img_np = np.zeros((64, 64, 3), dtype=np.uint8) - res = MagicMock() - res.fetch_file_data.return_value = img_np - return res - - -def test_predict_image_returns_box_annotations(module): - import numpy as np - fake_det = torch.tensor([[10.0, 20.0, 50.0, 60.0, 0.9, 0.85, 0.0]]) - with patch('yolox.utils.postprocess', return_value=[fake_det]): - results = module.predict_image([_make_fake_resource()]) - assert len(results) == 1 # one list per resource - assert len(results[0]) == 1 # one detection - assert isinstance(results[0][0], BoxAnnotation) - - -def test_predict_image_empty_returns_empty_list(module): - with patch('yolox.utils.postprocess', return_value=[None]): - results = module.predict_image([_make_fake_resource()]) - assert results == [[]] - - -def test_predict_image_box_coordinates(module): - fake_det = torch.tensor([[5.0, 10.0, 55.0, 70.0, 0.9, 0.9, 1.0]]) - with patch('yolox.utils.postprocess', return_value=[fake_det]): - results = module.predict_image([_make_fake_resource()]) - ann = results[0][0] - x1, y1, _ = ann.geometry.point1 - x2, y2, _ = ann.geometry.point2 - assert x1 == pytest.approx(5.0) - assert y1 == pytest.approx(10.0) - assert x2 == pytest.approx(55.0) - assert y2 == pytest.approx(70.0) - - -def test_predict_image_class_identifier(module): - fake_det = torch.tensor([[0.0, 0.0, 10.0, 10.0, 0.9, 0.9, 2.0]]) - with patch('yolox.utils.postprocess', return_value=[fake_det]): - results = module.predict_image([_make_fake_resource()]) - assert results[0][0].identifier == '2' - - -def test_predict_image_multiple_resources(module): - import numpy as np - fake_det = torch.tensor([[0.0, 0.0, 10.0, 10.0, 0.9, 0.9, 0.0]]) - resources = [_make_fake_resource(), _make_fake_resource()] - with patch('yolox.utils.postprocess', side_effect=[[fake_det], [None]]): - results = module.predict_image(resources) - assert len(results) == 2 - assert len(results[0]) == 1 - assert len(results[1]) == 0 diff --git a/tests/test_yolox_trainer.py b/tests/test_yolox_trainer.py index 72480f37..9b3c40a8 100644 --- a/tests/test_yolox_trainer.py +++ b/tests/test_yolox_trainer.py @@ -1,12 +1,14 @@ """Tests for YOLOXTrainer.""" -import pytest from unittest.mock import MagicMock, patch import albumentations as A +import pytest -from datamint.lightning.trainers.specialized.yolox import YOLOXTrainer from datamint.dataset.image_dataset import ImageDataset -from datamint.lightning.trainers.lightning_modules.detection_modules.yolox_module import YOLOXModule +from datamint.lightning.trainers.lightning_modules.detection_modules.yolox_module import ( + YOLOXModule, +) +from datamint.lightning.trainers.specialized.yolox import YOLOXTrainer @pytest.fixture()