diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 23ef380..fce6d8a 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -11,7 +11,7 @@ repos: args: ["--py313-plus"] files: ^itkit/ - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.21 + rev: v0.16.5 hooks: - id: ruff-check types_or: [ python, pyi ] @@ -36,7 +36,7 @@ repos: - id: requirements-txt-fixer - id: trailing-whitespace - repo: https://github.com/PyCQA/isort - rev: 9.0.0b1 + rev: 9.0.1 hooks: - id: isort files: ^itkit/ diff --git a/SlicerITKIT/ITKITInference/ITKITInference.py b/SlicerITKIT/ITKITInference/ITKITInference.py index d57eb85..07382df 100644 --- a/SlicerITKIT/ITKITInference/ITKITInference.py +++ b/SlicerITKIT/ITKITInference/ITKITInference.py @@ -20,7 +20,6 @@ import logging import os import tempfile -from typing import Optional import ctk import qt @@ -279,7 +278,6 @@ def setup(self) -> None: def cleanup(self) -> None: """Called when the application closes and the module widget is destroyed.""" - pass def onServerUrlChanged(self): """Called when server URL changes.""" @@ -330,8 +328,8 @@ def onConnectButton(self): LOGGER.exception("Connect to server failed") self.serverStatusLabel.setText("Connection failed") self.serverStatusLabel.setStyleSheet("color: red;") - self.statusLabel.setText(f"Error: {str(e)}") - slicer.util.errorDisplay(f"Failed to connect to server: {str(e)}") + self.statusLabel.setText(f"Error: {e!s}") + slicer.util.errorDisplay(f"Failed to connect to server: {e!s}") def updateModelStatus(self, model_info): """Update the current model status display.""" @@ -453,8 +451,8 @@ def onLoadModelButton(self): except Exception as e: LOGGER.exception("Load model failed") self.progressBar.hide() - self.statusLabel.setText(f"Error: {str(e)}") - slicer.util.errorDisplay(f"Failed to load model: {str(e)}") + self.statusLabel.setText(f"Error: {e!s}") + slicer.util.errorDisplay(f"Failed to load model: {e!s}") def onUnloadModelButton(self): """Unload the current model from the server.""" @@ -474,7 +472,7 @@ def onUnloadModelButton(self): except Exception as e: LOGGER.exception("Unload model failed") - slicer.util.errorDisplay(f"Failed to unload model: {str(e)}") + slicer.util.errorDisplay(f"Failed to unload model: {e!s}") def onApplyButton(self): """Run processing when user clicks "Apply" button.""" @@ -531,8 +529,8 @@ def onComplete(): LOGGER.exception("Inference failed") def onError(): - slicer.util.errorDisplay(f"Inference failed: {str(e)}") - self.statusLabel.setText(f"Error: {str(e)}") + slicer.util.errorDisplay(f"Inference failed: {e!s}") + self.statusLabel.setText(f"Error: {e!s}") self.applyButton.enabled = True self.progressBar.hide() @@ -572,7 +570,7 @@ def load_model( self, server_url: str, backend_type: str, - config_path: Optional[str], + config_path: str | None, model_path: str, inference_config: dict, ) -> bool: diff --git a/SlicerITKIT/server/itkit_server.py b/SlicerITKIT/server/itkit_server.py index b2f7ef0..205dcbd 100644 --- a/SlicerITKIT/server/itkit_server.py +++ b/SlicerITKIT/server/itkit_server.py @@ -26,6 +26,7 @@ import torch from flask import Flask, jsonify, request, send_file from flask_cors import CORS + from itkit.mm.inference import ( InferenceConfig, Inferencer_Seg3D, @@ -54,9 +55,9 @@ def __init__( self, name: str, backend_type: str, - config_path: Optional[str], + config_path: str | None, model_path: str, - inference_config: Optional[dict] = None, + inference_config: dict | None = None, ): self.name = name self.backend_type = backend_type # 'mmengine' or 'onnx' @@ -119,7 +120,7 @@ def to_dict(self): def _get_windowing_from_model( model: ModelConfig, -) -> tuple[Optional[float], Optional[float]]: +) -> tuple[float | None, float | None]: """Try to read window level/width from backend metadata or config.""" if model.backend_type.lower() == "mmengine": cfg = getattr(model.backend, "cfg", None) diff --git a/examples/configs/0.0.AbdomenCT1K_TorchIO/MedNeXt.py b/examples/configs/0.0.AbdomenCT1K_TorchIO/MedNeXt.py index d26f07f..d193667 100644 --- a/examples/configs/0.0.AbdomenCT1K_TorchIO/MedNeXt.py +++ b/examples/configs/0.0.AbdomenCT1K_TorchIO/MedNeXt.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.0.AbdomenCT1K_TorchIO/SegFormer3D.py b/examples/configs/0.0.AbdomenCT1K_TorchIO/SegFormer3D.py index 2fa14e0..0e77463 100644 --- a/examples/configs/0.0.AbdomenCT1K_TorchIO/SegFormer3D.py +++ b/examples/configs/0.0.AbdomenCT1K_TorchIO/SegFormer3D.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.0.AbdomenCT1K_TorchIO/mgam.py b/examples/configs/0.0.AbdomenCT1K_TorchIO/mgam.py index ae4fce1..e6925a3 100644 --- a/examples/configs/0.0.AbdomenCT1K_TorchIO/mgam.py +++ b/examples/configs/0.0.AbdomenCT1K_TorchIO/mgam.py @@ -1,35 +1,36 @@ -from torch.optim.adamw import AdamW -from torch.distributed.fsdp.api import ShardingStrategy - -from mmengine.runner import ValLoop -from mmengine.runner import TestLoop -from mmengine.hooks.iter_timer_hook import IterTimerHook -from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook -from mmengine.hooks.checkpoint_hook import CheckpointHook -from mmengine.hooks import DistSamplerSeedHook -from mmengine.runner import IterBasedTrainLoop -from mmengine.optim.scheduler import LinearLR, PolyLR -from mmengine.optim import OptimWrapper, AmpOptimWrapper -from mmengine.model.wrappers import MMFullyShardedDataParallel from mmengine._strategy.deepspeed import DeepSpeedOptimWrapper, DeepSpeedStrategy from mmengine.dataset.sampler import DefaultSampler, InfiniteSampler from mmengine.dataset.utils import default_collate +from mmengine.hooks import DistSamplerSeedHook +from mmengine.hooks.checkpoint_hook import CheckpointHook +from mmengine.hooks.iter_timer_hook import IterTimerHook +from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook +from mmengine.model.wrappers import MMFullyShardedDataParallel +from mmengine.optim import AmpOptimWrapper, OptimWrapper +from mmengine.optim.scheduler import LinearLR, PolyLR +from mmengine.runner import IterBasedTrainLoop, TestLoop, ValLoop from mmengine.visualization import TensorboardVisBackend +from torch.distributed.fsdp.api import ShardingStrategy +from torch.optim.adamw import AdamW + +from itkit.dataset import ITKITConcatDataset +from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha +from itkit.dataset.torchio import TorchIO_PatchedDataset # customize from itkit.mm.mmeng_PlugIn import ( - RemasteredDDP, LoggerJSON, RuntimeInfoHook, multi_sample_collate, - RatioSampler, RemasteredFSDP_Strategy) -from itkit.process.GeneralPreProcess import WindowSet, TypeConvert -from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA + LoggerJSON, + RatioSampler, + RemasteredDDP, + RemasteredFSDP_Strategy, + RuntimeInfoHook, + multi_sample_collate, +) from itkit.mm.mmseg_Dev3D import PackSeg3DInputs, Seg3DDataPreProcessor from itkit.mm.mmseg_PlugIn import IoUMetric_PerClass -from itkit.dataset import ITKITConcatDataset -from itkit.dataset.torchio import TorchIO_PatchedDataset -from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha -from itkit.mm.visualization import SegViser, BaseVisHook, LocalVisBackend - - +from itkit.mm.visualization import BaseVisHook, LocalVisBackend, SegViser +from itkit.process.GeneralPreProcess import TypeConvert, WindowSet +from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA # --------------------PARAMETERS-------------------- # diff --git a/examples/configs/0.1.AbdomenCT1K_MONAI/MedNeXt.py b/examples/configs/0.1.AbdomenCT1K_MONAI/MedNeXt.py index d26f07f..d193667 100644 --- a/examples/configs/0.1.AbdomenCT1K_MONAI/MedNeXt.py +++ b/examples/configs/0.1.AbdomenCT1K_MONAI/MedNeXt.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.1.AbdomenCT1K_MONAI/SegFormer3D.py b/examples/configs/0.1.AbdomenCT1K_MONAI/SegFormer3D.py index 2fa14e0..0e77463 100644 --- a/examples/configs/0.1.AbdomenCT1K_MONAI/SegFormer3D.py +++ b/examples/configs/0.1.AbdomenCT1K_MONAI/SegFormer3D.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.1.AbdomenCT1K_MONAI/mgam.py b/examples/configs/0.1.AbdomenCT1K_MONAI/mgam.py index db46ada..bb5bf99 100644 --- a/examples/configs/0.1.AbdomenCT1K_MONAI/mgam.py +++ b/examples/configs/0.1.AbdomenCT1K_MONAI/mgam.py @@ -1,35 +1,36 @@ -from torch.optim.adamw import AdamW -from torch.distributed.fsdp.api import ShardingStrategy - -from mmengine.runner import ValLoop -from mmengine.runner import TestLoop -from mmengine.hooks.iter_timer_hook import IterTimerHook -from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook -from mmengine.hooks.checkpoint_hook import CheckpointHook -from mmengine.hooks import DistSamplerSeedHook -from mmengine.runner import IterBasedTrainLoop -from mmengine.optim.scheduler import LinearLR, PolyLR -from mmengine.optim import OptimWrapper, AmpOptimWrapper -from mmengine.model.wrappers import MMFullyShardedDataParallel from mmengine._strategy.deepspeed import DeepSpeedOptimWrapper, DeepSpeedStrategy from mmengine.dataset.sampler import DefaultSampler, InfiniteSampler from mmengine.dataset.utils import default_collate +from mmengine.hooks import DistSamplerSeedHook +from mmengine.hooks.checkpoint_hook import CheckpointHook +from mmengine.hooks.iter_timer_hook import IterTimerHook +from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook +from mmengine.model.wrappers import MMFullyShardedDataParallel +from mmengine.optim import AmpOptimWrapper, OptimWrapper +from mmengine.optim.scheduler import LinearLR, PolyLR +from mmengine.runner import IterBasedTrainLoop, TestLoop, ValLoop from mmengine.visualization import TensorboardVisBackend +from torch.distributed.fsdp.api import ShardingStrategy +from torch.optim.adamw import AdamW + +from itkit.dataset import ITKITConcatDataset +from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha +from itkit.dataset.monai import MONAI_PatchedDataset # customize from itkit.mm.mmeng_PlugIn import ( - RemasteredDDP, LoggerJSON, RuntimeInfoHook, multi_sample_collate, - RatioSampler, RemasteredFSDP_Strategy) -from itkit.process.GeneralPreProcess import WindowSet, TypeConvert -from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA + LoggerJSON, + RatioSampler, + RemasteredDDP, + RemasteredFSDP_Strategy, + RuntimeInfoHook, + multi_sample_collate, +) from itkit.mm.mmseg_Dev3D import PackSeg3DInputs, Seg3DDataPreProcessor from itkit.mm.mmseg_PlugIn import IoUMetric_PerClass -from itkit.dataset import ITKITConcatDataset -from itkit.dataset.monai import MONAI_PatchedDataset -from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha -from itkit.mm.visualization import SegViser, BaseVisHook, LocalVisBackend - - +from itkit.mm.visualization import BaseVisHook, LocalVisBackend, SegViser +from itkit.process.GeneralPreProcess import TypeConvert, WindowSet +from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA # --------------------PARAMETERS-------------------- # diff --git a/examples/configs/0.2.AbdomenCT1K_ITKIT/MedNeXt.py b/examples/configs/0.2.AbdomenCT1K_ITKIT/MedNeXt.py index d26f07f..d193667 100644 --- a/examples/configs/0.2.AbdomenCT1K_ITKIT/MedNeXt.py +++ b/examples/configs/0.2.AbdomenCT1K_ITKIT/MedNeXt.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.2.AbdomenCT1K_ITKIT/SegFormer3D.py b/examples/configs/0.2.AbdomenCT1K_ITKIT/SegFormer3D.py index 2fa14e0..0e77463 100644 --- a/examples/configs/0.2.AbdomenCT1K_ITKIT/SegFormer3D.py +++ b/examples/configs/0.2.AbdomenCT1K_ITKIT/SegFormer3D.py @@ -1,4 +1,5 @@ from mmengine.config import read_base + with read_base(): from .mgam import * diff --git a/examples/configs/0.2.AbdomenCT1K_ITKIT/mgam.py b/examples/configs/0.2.AbdomenCT1K_ITKIT/mgam.py index e466d64..23597e7 100644 --- a/examples/configs/0.2.AbdomenCT1K_ITKIT/mgam.py +++ b/examples/configs/0.2.AbdomenCT1K_ITKIT/mgam.py @@ -1,34 +1,35 @@ -from torch.optim.adamw import AdamW -from torch.distributed.fsdp.api import ShardingStrategy - -from mmengine.runner import ValLoop -from mmengine.runner import TestLoop -from mmengine.hooks.iter_timer_hook import IterTimerHook -from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook -from mmengine.hooks.checkpoint_hook import CheckpointHook -from mmengine.hooks import DistSamplerSeedHook -from mmengine.runner import IterBasedTrainLoop -from mmengine.optim.scheduler import LinearLR, PolyLR -from mmengine.optim import OptimWrapper, AmpOptimWrapper -from mmengine.model.wrappers import MMFullyShardedDataParallel from mmengine._strategy.deepspeed import DeepSpeedOptimWrapper, DeepSpeedStrategy from mmengine.dataset.sampler import DefaultSampler, InfiniteSampler from mmengine.dataset.utils import default_collate +from mmengine.hooks import DistSamplerSeedHook +from mmengine.hooks.checkpoint_hook import CheckpointHook +from mmengine.hooks.iter_timer_hook import IterTimerHook +from mmengine.hooks.param_scheduler_hook import ParamSchedulerHook +from mmengine.model.wrappers import MMFullyShardedDataParallel +from mmengine.optim import AmpOptimWrapper, OptimWrapper +from mmengine.optim.scheduler import LinearLR, PolyLR +from mmengine.runner import IterBasedTrainLoop, TestLoop, ValLoop from mmengine.visualization import TensorboardVisBackend +from torch.distributed.fsdp.api import ShardingStrategy +from torch.optim.adamw import AdamW + +from itkit.dataset import ITKITConcatDataset +from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha, AbdomenCT_1K_Patch # customize from itkit.mm.mmeng_PlugIn import ( - RemasteredDDP, LoggerJSON, RuntimeInfoHook, multi_sample_collate, - RatioSampler, RemasteredFSDP_Strategy) -from itkit.process.GeneralPreProcess import WindowSet, TypeConvert -from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA + LoggerJSON, + RatioSampler, + RemasteredDDP, + RemasteredFSDP_Strategy, + RuntimeInfoHook, + multi_sample_collate, +) from itkit.mm.mmseg_Dev3D import PackSeg3DInputs, Seg3DDataPreProcessor from itkit.mm.mmseg_PlugIn import IoUMetric_PerClass -from itkit.dataset import ITKITConcatDataset -from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha, AbdomenCT_1K_Patch -from itkit.mm.visualization import SegViser, BaseVisHook, LocalVisBackend - - +from itkit.mm.visualization import BaseVisHook, LocalVisBackend, SegViser +from itkit.process.GeneralPreProcess import TypeConvert, WindowSet +from itkit.process.LoadBiomedicalData import LoadImageFromMHA, LoadMaskFromMHA # --------------------PARAMETERS-------------------- # diff --git a/itkit/dataset/BraTs2024/convert_nii_mha.py b/itkit/dataset/BraTs2024/convert_nii_mha.py index 1f5cc17..a52616e 100644 --- a/itkit/dataset/BraTs2024/convert_nii_mha.py +++ b/itkit/dataset/BraTs2024/convert_nii_mha.py @@ -131,15 +131,14 @@ def convert_brats_to_mha(input_dir, dest_root, spacing=None, size=None, use_mp=F partial_convert_func = partial(convert_case, dest_root=dest_root, spacing=spacing, size=size) if use_mp: - with mp.Pool(mp.cpu_count()) as pool: - with tqdm( - total=len(case_dirs), - desc="Converting BraTs2024", - dynamic_ncols=True, - leave=False - ) as pbar: - for _ in pool.imap_unordered(partial_convert_func, case_dirs): - pbar.update() + with mp.Pool(mp.cpu_count()) as pool, tqdm( + total=len(case_dirs), + desc="Converting BraTs2024", + dynamic_ncols=True, + leave=False + ) as pbar: + for _ in pool.imap_unordered(partial_convert_func, case_dirs): + pbar.update() else: with tqdm(total=len(case_dirs), desc="Converting Cases") as pbar: for case_dir in case_dirs: diff --git a/itkit/dataset/Totalsegmentator/fix_ITKUnreadable_files.py b/itkit/dataset/Totalsegmentator/fix_ITKUnreadable_files.py index 4f31c21..363f587 100644 --- a/itkit/dataset/Totalsegmentator/fix_ITKUnreadable_files.py +++ b/itkit/dataset/Totalsegmentator/fix_ITKUnreadable_files.py @@ -19,7 +19,7 @@ def fix_nifti_file(input_path, output_path): nib.save(img, output_path) return True except Exception as e: - print(f"处理文件 {input_path} 时出错: {str(e)}") + print(f"处理文件 {input_path} 时出错: {e!s}") return False @@ -39,7 +39,7 @@ def process_file(file_path, input_root, output_root): output_path = output_root / rel_path return fix_nifti_file(file_path, output_path) except ValueError as e: - print(f"计算相对路径时出错: {str(e)}") + print(f"计算相对路径时出错: {e!s}") return False diff --git a/itkit/dataset/base.py b/itkit/dataset/base.py index ab34bff..4a02339 100644 --- a/itkit/dataset/base.py +++ b/itkit/dataset/base.py @@ -223,7 +223,7 @@ def sample_iterator(self) -> Generator[tuple[str, str]]: if not os.path.exists(image_mha_path): print_log(f"{series} image mha file not found.\nFullPath: {image_mha_path}", MMLogger.get_current_instance(), - logging.WARN) + logging.WARNING) continue yield (image_mha_path, label_mha_path) @@ -253,11 +253,11 @@ def sample_iterator(self) -> Generator[tuple[str, str]]: # List all image files that match the current series UID # Files are in format: _.mha (e.g., 1.3.6.1.4.1.9328.50.4.0095_p0.mha) if not os.path.exists(image_folder): - print_log(f"Image folder not found: {image_folder}", MMLogger.get_current_instance(), logging.WARN) + print_log(f"Image folder not found: {image_folder}", MMLogger.get_current_instance(), logging.WARNING) continue if series not in self.precrop_meta["patch_meta"]: - print_log(f"Series {series} not found in patch metadata", MMLogger.get_current_instance(), logging.WARN) + print_log(f"Series {series} not found in patch metadata", MMLogger.get_current_instance(), logging.WARNING) continue series_patch_files = self.precrop_meta["patch_meta"][series]["class_within_patch"].keys() diff --git a/itkit/io/nii_toolkit.py b/itkit/io/nii_toolkit.py index deb5827..c08596c 100644 --- a/itkit/io/nii_toolkit.py +++ b/itkit/io/nii_toolkit.py @@ -14,7 +14,7 @@ def convert_nii_sitk(nii_path:str, nii_fdata_order:Literal['xyz','zyx'], dtype=np.float32, - value_offset:int|float|None=None + value_offset:float | None=None ) -> sitk.Image: # 加载并进行值域修正 try: diff --git a/itkit/io/sitk_toolkit.py b/itkit/io/sitk_toolkit.py index 4d1298e..8f494b2 100644 --- a/itkit/io/sitk_toolkit.py +++ b/itkit/io/sitk_toolkit.py @@ -159,7 +159,7 @@ def sitk_new_blank_image(size, spacing, direction, origin, default_value=0.0): def nii_to_sitk( nii_path: str, field: Literal["image", "label"], - value_offset: int | float | None = None, + value_offset: float | None = None, ) -> sitk.Image: try: sitk_img = sitk.ReadImage(nii_path, outputPixelType=sitk.sitkInt16 if field == "image" else sitk.sitkUInt8) diff --git a/itkit/mm/inference.py b/itkit/mm/inference.py index 3585391..9ac0291 100644 --- a/itkit/mm/inference.py +++ b/itkit/mm/inference.py @@ -44,7 +44,6 @@ def forward(self, inputs: Tensor) -> Tensor: Returns: Output logits tensor. """ - pass @abstractmethod def slide_inference(self, inputs: Tensor) -> Tensor: @@ -56,7 +55,6 @@ def slide_inference(self, inputs: Tensor) -> Tensor: Returns: Segmentation logits tensor. """ - pass class MMEngineInferBackend(InferenceBackend): @@ -270,7 +268,6 @@ def Inference_FromNDArray(self, inputs: np.ndarray) -> tuple[Tensor, Tensor]: - seg_logits (Tensor): Segmentation logits tensor. - sem_seg_map (Tensor): Segmentation map tensor. """ - pass class Inferencer_Seg3D(Inferencer): diff --git a/itkit/mm/mmeng_PlugIn.py b/itkit/mm/mmeng_PlugIn.py index 60d0bbd..eb30d52 100644 --- a/itkit/mm/mmeng_PlugIn.py +++ b/itkit/mm/mmeng_PlugIn.py @@ -73,7 +73,7 @@ def custom_env(self, cfg): # Avoid device clash with OpenCV torch.cuda.set_device(cfg.pop("torch_cuda_id", -1)) # Torch Compile - cfg.get("torch_logging_level", logging.WARN) + cfg.get("torch_logging_level", logging.WARNING) torch._logging.set_logs(all=self.str_to_log_level(cfg.pop("torch_logging_level", "WARN")), dynamo=self.str_to_log_level(cfg.pop("dynamo_logging_level", "WARN"))) torch._dynamo.config.cache_size_limit = cfg.pop("dynamo_cache_size", 1) @@ -102,7 +102,7 @@ def auto_configure_num_classes_from_Databackend(cfg: ConfigType, num_classes): def load_or_resume(self) -> None: if self._has_loaded: - return None + return # Resume has higher priority than `load_from` if self._resume: diff --git a/itkit/mm/mmseg_Dev3D.py b/itkit/mm/mmseg_Dev3D.py index 7bfcb10..562e1c4 100644 --- a/itkit/mm/mmseg_Dev3D.py +++ b/itkit/mm/mmseg_Dev3D.py @@ -16,7 +16,10 @@ from mmseg.models.decode_heads.decode_head import BaseDecodeHead from mmseg.models.losses.accuracy import accuracy from mmseg.models.segmentors.encoder_decoder import EncoderDecoder -from mmseg.structures.seg_data_sample import PixelData, SegDataSample # pyright: ignore[reportAttributeAccessIssue] +from mmseg.structures.seg_data_sample import ( # pyright: ignore[reportAttributeAccessIssue] + PixelData, + SegDataSample, +) from mmseg.visualization.local_visualizer import SegLocalVisualizer from torch import Tensor from torch.nn import functional as F @@ -395,8 +398,8 @@ def postprocess_result( i_seg_pred = (i_seg_logits > self.decode_head.threshold).to(i_seg_logits) data_samples[i].set_data( { - "seg_logits": VolumeData(**{"data": i_seg_logits}), # type: ignore - "pred_sem_seg": VolumeData(**{"data": i_seg_pred}), # type: ignore + "seg_logits": VolumeData(data=i_seg_logits), # type: ignore + "pred_sem_seg": VolumeData(data=i_seg_pred), # type: ignore } ) @@ -848,8 +851,8 @@ def __init__( std: Sequence[float] | None = None, size: tuple | None = None, size_divisor: int | None = None, - pad_val: int | float = 0, - seg_pad_val: int | float = 255, + pad_val: float = 0, + seg_pad_val: float = 255, rot3D_angle: Sequence | None = None, test_cfg: dict | None = None, non_blocking: bool = True, @@ -882,8 +885,8 @@ def stack_batch_3D( data_samples: list[Seg3DDataSample], size: tuple | None = None, size_divisor: int | None = None, - pad_val: int | float = 0, - seg_pad_val: int | float = 255, + pad_val: float = 0, + seg_pad_val: float = 255, training: bool = True, ): """Stack multiple 3D volume inputs to form a batch and pad the volumes and gt_sem_segs diff --git a/itkit/mm/run.py b/itkit/mm/run.py index f9f8ba9..bda3ba3 100644 --- a/itkit/mm/run.py +++ b/itkit/mm/run.py @@ -99,9 +99,8 @@ def find_full_exp_name(self, exp_name): print(f"Found experiment by prefix: {exp_name} -> {exp}") return exp - else: - print(f"No experiment found under {self.config_root} directory: {exp_name}") - return None + print(f"No experiment found under {self.config_root} directory: {exp_name}") + return None def experiment_queue(self): print("Experiment queue started, importing dependencies...") diff --git a/itkit/mm/task_models.py b/itkit/mm/task_models.py index 4f876a7..7839648 100644 --- a/itkit/mm/task_models.py +++ b/itkit/mm/task_models.py @@ -395,8 +395,8 @@ def _predict(force_cpu:bool=False): i_seg_pred = (i_seg_logits_sigmoid > self.binary_segment_threshold).to(i_seg_logits) # Store results into data_samples - data_samples[i].seg_logits = VolumeData(**{"data": i_seg_logits}) # pyright: ignore[reportArgumentType] - data_samples[i].pred_sem_seg = VolumeData(**{"data": i_seg_pred}) # pyright: ignore[reportArgumentType] + data_samples[i].seg_logits = VolumeData(data=i_seg_logits) # pyright: ignore[reportArgumentType] + data_samples[i].pred_sem_seg = VolumeData(data=i_seg_pred) # pyright: ignore[reportArgumentType] return data_samples diff --git a/itkit/mm/visualization.py b/itkit/mm/visualization.py index 0e29423..82cb674 100644 --- a/itkit/mm/visualization.py +++ b/itkit/mm/visualization.py @@ -317,7 +317,7 @@ def add_datasample(self, if gt_seg_map is None: print_log(f"When visualizing `{name}` with img_path `{img_path}`, " "gt_seg_map is None. So the gt_seg_map will not be empty.", - MMLogger.get_current_instance(), logging.WARN) + MMLogger.get_current_instance(), logging.WARNING) # draw fig and save image_array = self._draw_fig(img_path, image_cpu, gt_seg_map, pred_seg_map, pred_seg_logits) diff --git a/itkit/models/DA_TransUnet/DATransUNet.py b/itkit/models/DA_TransUnet/DATransUNet.py index b72170c..9308213 100644 --- a/itkit/models/DA_TransUnet/DATransUNet.py +++ b/itkit/models/DA_TransUnet/DATransUNet.py @@ -7,8 +7,8 @@ import numpy as np import torch -import torch.nn as nn from scipy import ndimage +from torch import nn from torch.nn import ( Conv2d, Dropout, diff --git a/itkit/models/DconnNet/DconnNet.py b/itkit/models/DconnNet/DconnNet.py index ed92246..f19cebf 100644 --- a/itkit/models/DconnNet/DconnNet.py +++ b/itkit/models/DconnNet/DconnNet.py @@ -6,7 +6,7 @@ import math import torch -import torch.nn as nn +from torch import nn # from resnet import resnet34 # import resnet diff --git a/itkit/models/DconnNet/gap.py b/itkit/models/DconnNet/gap.py index 8685034..683eea0 100644 --- a/itkit/models/DconnNet/gap.py +++ b/itkit/models/DconnNet/gap.py @@ -1,5 +1,5 @@ import torch -import torch.nn as nn +from torch import nn GlobalAvgPool2D = lambda: nn.AdaptiveAvgPool2d(1) diff --git a/itkit/models/DconnNet/resnet.py b/itkit/models/DconnNet/resnet.py index f527e2f..21e2f34 100644 --- a/itkit/models/DconnNet/resnet.py +++ b/itkit/models/DconnNet/resnet.py @@ -6,8 +6,8 @@ import math -import torch.nn as nn -import torch.utils.model_zoo as model_zoo +from torch import nn +from torch.utils import model_zoo __all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101', 'resnet152'] diff --git a/itkit/models/EfficientFormer.py b/itkit/models/EfficientFormer.py index c72400b..6043335 100644 --- a/itkit/models/EfficientFormer.py +++ b/itkit/models/EfficientFormer.py @@ -3,8 +3,8 @@ import timm import torch -import torch.nn as nn import torch.nn.functional as F +from torch import nn class EfficientFormerV2(torch.nn.Module): diff --git a/itkit/models/EfficientNet.py b/itkit/models/EfficientNet.py index e7ef740..6086f59 100644 --- a/itkit/models/EfficientNet.py +++ b/itkit/models/EfficientNet.py @@ -1,8 +1,8 @@ # pyright: reportCallIssue=false import timm import torch -import torch.nn as nn import torch.nn.functional as F +from torch import nn class EfficientNetV2(torch.nn.Module): diff --git a/itkit/models/LM_Net/FrequencyDomain.py b/itkit/models/LM_Net/FrequencyDomain.py index 281f412..dfe3b68 100644 --- a/itkit/models/LM_Net/FrequencyDomain.py +++ b/itkit/models/LM_Net/FrequencyDomain.py @@ -1,6 +1,6 @@ import torch -import torch.nn as nn from resnet import resnet50 +from torch import nn #from torchvision.models import resnet50 diff --git a/itkit/models/LM_Net/LM_Net.py b/itkit/models/LM_Net/LM_Net.py index a0a3f99..44de7a2 100644 --- a/itkit/models/LM_Net/LM_Net.py +++ b/itkit/models/LM_Net/LM_Net.py @@ -1,6 +1,6 @@ # pyright: reportCallIssue=false -import torch.nn as nn +from torch import nn #from .nonlocal_block import NONLocalBlock2D from .modules import * diff --git a/itkit/models/LM_Net/acb.py b/itkit/models/LM_Net/acb.py index 93e0203..c3b6415 100644 --- a/itkit/models/LM_Net/acb.py +++ b/itkit/models/LM_Net/acb.py @@ -1,8 +1,8 @@ from typing import Literal import torch -import torch.nn as nn -import torch.nn.init as init +from torch import nn +from torch.nn import init class ACBlock(nn.Module): diff --git a/itkit/models/LM_Net/blur_pool.py b/itkit/models/LM_Net/blur_pool.py index 0b9619d..059a95e 100644 --- a/itkit/models/LM_Net/blur_pool.py +++ b/itkit/models/LM_Net/blur_pool.py @@ -8,8 +8,8 @@ import numpy as np import torch -import torch.nn as nn import torch.nn.functional as F +from torch import nn class BlurPool2d(nn.Module): diff --git a/itkit/models/LM_Net/depthwise_conv2d_implicit_gemm.py b/itkit/models/LM_Net/depthwise_conv2d_implicit_gemm.py index 05df3bd..0f3324d 100644 --- a/itkit/models/LM_Net/depthwise_conv2d_implicit_gemm.py +++ b/itkit/models/LM_Net/depthwise_conv2d_implicit_gemm.py @@ -2,8 +2,8 @@ import _depthwise_conv2d_implicit_gemm_C as _extension import torch -import torch.nn as nn from depthwise_conv2d_implicit_gemm import * +from torch import nn class _DepthWiseConv2dImplicitGEMMFP32(torch.autograd.Function): diff --git a/itkit/models/LM_Net/involution_cuda.py b/itkit/models/LM_Net/involution_cuda.py index 462d205..2606bff 100644 --- a/itkit/models/LM_Net/involution_cuda.py +++ b/itkit/models/LM_Net/involution_cuda.py @@ -3,7 +3,7 @@ import cupy import torch -import torch.nn as nn +from torch import nn from torch.autograd import Function from torch.nn.modules.utils import _pair diff --git a/itkit/models/LM_Net/involution_naive.py b/itkit/models/LM_Net/involution_naive.py index 7181a9b..da9107c 100644 --- a/itkit/models/LM_Net/involution_naive.py +++ b/itkit/models/LM_Net/involution_naive.py @@ -1,4 +1,4 @@ -import torch.nn as nn +from torch import nn class involution(nn.Module): diff --git a/itkit/models/LM_Net/modules.py b/itkit/models/LM_Net/modules.py index 30beee8..c871563 100644 --- a/itkit/models/LM_Net/modules.py +++ b/itkit/models/LM_Net/modules.py @@ -2,13 +2,13 @@ from collections import OrderedDict import torch -import torch.nn as nn import torch.nn.functional as F #from depthwise_conv2d_implicit_gemm import DepthWiseConv2dImplicitGEMM #from .involution_cuda import involution from natten import NeighborhoodAttention2D from timm.models.layers import DropPath, to_2tuple, trunc_normal_ +from torch import nn # from .nonlocal_block import NONLocalBlock2D #from carafe import CARAFEPack @@ -17,7 +17,10 @@ # Try to import optional dependencies try: - from .nattencuda import NeighborhoodAttention, NEWNeighborhoodAttention # type: ignore + from .nattencuda import ( # type: ignore + NeighborhoodAttention, + NEWNeighborhoodAttention, + ) except ImportError: NEWNeighborhoodAttention = None # type: ignore NeighborhoodAttention = None # type: ignore diff --git a/itkit/models/LM_Net/resnet.py b/itkit/models/LM_Net/resnet.py index b0faf40..c0d7a98 100644 --- a/itkit/models/LM_Net/resnet.py +++ b/itkit/models/LM_Net/resnet.py @@ -1,4 +1,4 @@ -import torch.nn as nn +from torch import nn from torch.hub import load_state_dict_from_url from torchvision.models import resnet50 @@ -132,7 +132,7 @@ def __init__(self, block, layers, num_classes=1000, zero_init_residual=False, replace_stride_with_dilation = [False, False, False] if len(replace_stride_with_dilation) != 3: raise ValueError("replace_stride_with_dilation should be None " - "or a 3-element tuple, got {}".format(replace_stride_with_dilation)) + f"or a 3-element tuple, got {replace_stride_with_dilation}") self.groups = groups self.base_width = width_per_group self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=2, padding=3, diff --git a/itkit/models/SegMamba.py b/itkit/models/SegMamba.py index 69fa31b..f06ba14 100644 --- a/itkit/models/SegMamba.py +++ b/itkit/models/SegMamba.py @@ -22,11 +22,11 @@ from __future__ import annotations import torch -import torch.nn as nn import torch.nn.functional as F from mamba_ssm import Mamba from monai.networks.blocks.dynunet_block import UnetOutBlock from monai.networks.blocks.unetr_block import UnetrBasicBlock, UnetrUpBlock +from torch import nn class LayerNorm(nn.Module): diff --git a/itkit/models/SwinUMamba/SwinUMamba.py b/itkit/models/SwinUMamba/SwinUMamba.py index 136f349..5fff155 100644 --- a/itkit/models/SwinUMamba/SwinUMamba.py +++ b/itkit/models/SwinUMamba/SwinUMamba.py @@ -4,12 +4,12 @@ from functools import partial import torch -import torch.nn as nn import torch.nn.functional as F -import torch.utils.checkpoint as checkpoint from einops import rearrange, repeat from mamba_ssm.ops.selective_scan_interface import selective_scan_fn from timm.models.layers import DropPath, trunc_normal_ +from torch import nn +from torch.utils import checkpoint DropPath.__repr__ = lambda self: f"timm.DropPath({self.drop_prob})" diff --git a/itkit/models/UNETR.py b/itkit/models/UNETR.py index 7bdb804..5a6821f 100644 --- a/itkit/models/UNETR.py +++ b/itkit/models/UNETR.py @@ -12,11 +12,11 @@ from collections.abc import Sequence import torch -import torch.nn as nn from monai.networks.blocks import UnetrBasicBlock, UnetrPrUpBlock, UnetrUpBlock from monai.networks.blocks.dynunet_block import UnetOutBlock from monai.networks.blocks.patchembedding import PatchEmbeddingBlock from monai.networks.blocks.transformerblock import TransformerBlock +from torch import nn class ViT(nn.Module): diff --git a/itkit/models/UNet3Plus.py b/itkit/models/UNet3Plus.py index ea57a25..33ccfff 100644 --- a/itkit/models/UNet3Plus.py +++ b/itkit/models/UNet3Plus.py @@ -1,8 +1,8 @@ from collections.abc import Sequence import torch -import torch.nn as nn import torch.nn.functional as F +from torch import nn from torch.utils.checkpoint import checkpoint diff --git a/itkit/models/VMamba/volume_mamba.py b/itkit/models/VMamba/volume_mamba.py index 69c2340..08d81e0 100644 --- a/itkit/models/VMamba/volume_mamba.py +++ b/itkit/models/VMamba/volume_mamba.py @@ -13,9 +13,9 @@ import math import torch -import torch.nn as nn import torch.utils.checkpoint from selective_scan import selective_scan_fn +from torch import nn from vmamba import mamba_init # pyright: ignore[reportMissingImports] diff --git a/itkit/models/mednext/MedNextV1.py b/itkit/models/mednext/MedNextV1.py index 61db2c0..4b6c513 100644 --- a/itkit/models/mednext/MedNextV1.py +++ b/itkit/models/mednext/MedNextV1.py @@ -1,9 +1,8 @@ from typing import cast import torch -import torch.nn as nn -import torch.utils.checkpoint as checkpoint -from torch import Tensor +from torch import Tensor, nn +from torch.utils import checkpoint from .blocks import * diff --git a/itkit/models/mednext/blocks.py b/itkit/models/mednext/blocks.py index 835fc9d..1b811cb 100644 --- a/itkit/models/mednext/blocks.py +++ b/itkit/models/mednext/blocks.py @@ -1,6 +1,6 @@ import torch -import torch.nn as nn import torch.nn.functional as F +from torch import nn class MedNeXtBlock(nn.Module): diff --git a/itkit/process/GeneralPreProcess.py b/itkit/process/GeneralPreProcess.py index 34269ce..a4b21a3 100644 --- a/itkit/process/GeneralPreProcess.py +++ b/itkit/process/GeneralPreProcess.py @@ -514,13 +514,12 @@ def generate_crop_bbox(img: np.ndarray) -> tuple: # when pass all check return crop_bbox - else: - raise RuntimeError( - Fore.YELLOW + \ + raise RuntimeError( + Fore.YELLOW + \ f"Cannot find a valid crop bbox after {self.CROP_RETRY+1} trials. " + \ f"Last check result: ccm_check={ccm_check_}, std_check={std_check_}." + \ Style.RESET_ALL - ) + ) def crop(self, img: np.ndarray, crop_bbox: tuple) -> np.ndarray: """Crop from ``img`` @@ -599,7 +598,7 @@ class RandomContinuousErase(BaseTransform): def __init__( self, max_size: list[int] | int, - pad_val: float | int, + pad_val: float, seg_pad_val=0, prob: float = 0.5, ): @@ -1025,7 +1024,12 @@ def __init__(self, method: Literal['gibbs', 'gaussian', 'kspace', 'rician'], prob: float = 0.5, **kwargs): - from monai.transforms import RandGaussianNoise, RandGibbsNoise, RandKSpaceSpikeNoise, RandRicianNoise + from monai.transforms import ( + RandGaussianNoise, + RandGibbsNoise, + RandKSpaceSpikeNoise, + RandRicianNoise, + ) if method == 'gibbs': self.noise_fn = RandGibbsNoise(prob, **kwargs) elif method == 'gaussian': diff --git a/itkit/process/base_processor.py b/itkit/process/base_processor.py index 5f46023..315cbad 100644 --- a/itkit/process/base_processor.py +++ b/itkit/process/base_processor.py @@ -186,7 +186,6 @@ def generate_metadata_for_existing_files(self): Subclasses should override this if they skip existing files. """ - pass def _generate_metadata_for_folder(self, dest_folder: str, source_folder: str | None, source_files_set: set[str] | None = None) -> None: @@ -244,8 +243,7 @@ def _generate_metadata_for_folder(self, dest_folder: str, source_folder: str | N def _normalize_filename(self, filepath: str) -> str: base = os.path.splitext(filepath)[0] # Handle double extensions like .nii.gz - if base.endswith('.nii'): - base = base[:-4] + base = base.removesuffix('.nii') return base def _collect_results(self, results: list): diff --git a/itkit/process/itk_check.py b/itkit/process/itk_check.py index 054b7a4..41127cc 100644 --- a/itkit/process/itk_check.py +++ b/itkit/process/itk_check.py @@ -223,7 +223,7 @@ def process_one(self, args: tuple[str, str]) -> tuple[SeriesMetadata | None, Val return SeriesMetadata.from_sitk_image(lbl, name), res except Exception as e: - res = ValidationResult(name, False, [f"Failed to read: {str(e)}"], (img_path, lbl_path)) + res = ValidationResult(name, False, [f"Failed to read: {e!s}"], (img_path, lbl_path)) return None, res def _collect_results(self, results: list): @@ -326,7 +326,7 @@ def process_one(self, args) -> tuple[SeriesMetadata | None, ValidationResult]: return SeriesMetadata.from_sitk_image(img, name), res except Exception as e: - res = ValidationResult(name, False, [f"Failed to read: {str(e)}"], img_path) + res = ValidationResult(name, False, [f"Failed to read: {e!s}"], img_path) return None, res def _collect_results(self, results: list): diff --git a/itkit/process/itk_extract.py b/itkit/process/itk_extract.py index 3a54b68..a01d1ea 100644 --- a/itkit/process/itk_extract.py +++ b/itkit/process/itk_extract.py @@ -37,8 +37,7 @@ def process_one(self, args: str) -> SeriesMetadata | None: # Normalize extension to .mha base_name = os.path.splitext(os.path.basename(output_path))[0] - if base_name.endswith('.nii'): - base_name = base_name[:-4] + base_name = base_name.removesuffix('.nii') output_path = os.path.join(os.path.dirname(output_path), base_name + '.mha') return self._extract_one_sample(file_path, output_path) diff --git a/itkit/process/itk_infer.py b/itkit/process/itk_infer.py index 28a6268..9dbb42c 100644 --- a/itkit/process/itk_infer.py +++ b/itkit/process/itk_infer.py @@ -91,7 +91,11 @@ def process_gpu_task(process_id, file_list, args, pred_conf_shared_dict=None): # NOTE Local environment setup for each GPU process. gpu_id = process_id % args.gpus os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id) - from itkit.mm.inference import Inferencer_Seg3D, MMEngineInferBackend, ONNXInferBackend + from itkit.mm.inference import ( + Inferencer_Seg3D, + MMEngineInferBackend, + ONNXInferBackend, + ) from itkit.mm.sliding_window import InferenceConfig tqdm.write(f"Process {process_id} using GPU {gpu_id}, processing {len(file_list)} files") diff --git a/itkit/web/app.py b/itkit/web/app.py index 31ea687..61e17ae 100644 --- a/itkit/web/app.py +++ b/itkit/web/app.py @@ -18,7 +18,14 @@ import uuid from pathlib import Path -from flask import Flask, Response, jsonify, render_template, request, stream_with_context +from flask import ( + Flask, + Response, + jsonify, + render_template, + request, + stream_with_context, +) app = Flask(__name__) diff --git a/tests/conftest.py b/tests/conftest.py index 40f9113..4e91c62 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,6 +1,6 @@ +import json import os import random -import json import numpy as np import pytest diff --git a/tests/dataset/test_dataset_extended.py b/tests/dataset/test_dataset_extended.py index e07f607..29cdf71 100644 --- a/tests/dataset/test_dataset_extended.py +++ b/tests/dataset/test_dataset_extended.py @@ -1,9 +1,11 @@ from unittest.mock import MagicMock, patch +from mmengine.logging import MMLogger + from itkit.dataset.base import SeriesVolumeDataset from itkit.dataset.monai import MONAI_PatchedDataset from itkit.dataset.torchio import TorchIO_PatchedDataset -from mmengine.logging import MMLogger + class SimpleSeriesDataset(SeriesVolumeDataset): """A minimal concrete implementation for testing split logic""" diff --git a/tests/dataset/test_dataset_filtering.py b/tests/dataset/test_dataset_filtering.py index d25fe51..758a1c5 100644 --- a/tests/dataset/test_dataset_filtering.py +++ b/tests/dataset/test_dataset_filtering.py @@ -1,11 +1,14 @@ -import os import json -import pytest +import os import tempfile + +import pytest import SimpleITK as sitk -from itkit.dataset.base import SeriesVolumeDataset from mmengine.logging import MMLogger +from itkit.dataset.base import SeriesVolumeDataset + + def create_test_image(path: str, size: tuple, spacing: tuple): """Helper to create test MHA images (Size and Spacing in Z, Y, X)""" # SimpleITK uses XYZ ordering, so we reverse ZYX to XYZ diff --git a/tests/dataset/test_dataset_registry.py b/tests/dataset/test_dataset_registry.py index eff2f7a..9118d1c 100644 --- a/tests/dataset/test_dataset_registry.py +++ b/tests/dataset/test_dataset_registry.py @@ -1,17 +1,19 @@ -import pytest from unittest.mock import patch + +import pytest from mmengine.logging import MMLogger # Import dataset classes from itkit.dataset.AbdomenCT_1K.mm_dataset import AbdomenCT_1K_Mha from itkit.dataset.CT_ORG.mm_dataset import CT_ORG_Mha -from itkit.dataset.KiTS23.mm_dataset import KiTS23_Mha +from itkit.dataset.CTSpine1K.mm_dataset import CTSpine1K_Mha from itkit.dataset.FLARE_2022.mm_dataset import FLARE_2022_Mha from itkit.dataset.FLARE_2023.mm_dataset import FLARE_2023_Mha -from itkit.dataset.CTSpine1K.mm_dataset import CTSpine1K_Mha from itkit.dataset.ImageTBAD.mm_dataset import TBAD_Mha -from itkit.dataset.LUNA16.mm_dataset import LUNA16_Mha +from itkit.dataset.KiTS23.mm_dataset import KiTS23_Mha from itkit.dataset.LiTS.mm_dataset import LiTS_Mha +from itkit.dataset.LUNA16.mm_dataset import LUNA16_Mha + @pytest.mark.parametrize("dataset_class, extra_kwargs", [ (AbdomenCT_1K_Mha, {}), diff --git a/tests/dataset/test_monai_integration.py b/tests/dataset/test_monai_integration.py index f1120ca..135949b 100644 --- a/tests/dataset/test_monai_integration.py +++ b/tests/dataset/test_monai_integration.py @@ -1,9 +1,11 @@ import os import shutil import tempfile + import numpy as np import pytest import SimpleITK as sitk + from itkit.dataset.monai import MONAI_PatchedDataset diff --git a/tests/dataset/test_torchio_integration.py b/tests/dataset/test_torchio_integration.py index c8c0cb1..c29ee92 100644 --- a/tests/dataset/test_torchio_integration.py +++ b/tests/dataset/test_torchio_integration.py @@ -1,10 +1,12 @@ import os import shutil import tempfile + import numpy as np import pytest import SimpleITK as sitk import torchio as tio + from itkit.dataset.torchio import TorchIO_PatchedDataset diff --git a/tests/itk_process/test_itk_convert.py b/tests/itk_process/test_itk_convert.py index 38c1314..03408d1 100644 --- a/tests/itk_process/test_itk_convert.py +++ b/tests/itk_process/test_itk_convert.py @@ -1,10 +1,10 @@ """Tests for itk_convert module - ITKIT to MONAI and TorchIO format conversion.""" import csv +import importlib.util import json import os import tempfile -import importlib.util import numpy as np import pytest diff --git a/tests/itk_process/test_process_load.py b/tests/itk_process/test_process_load.py index 685e4d9..d16f519 100644 --- a/tests/itk_process/test_process_load.py +++ b/tests/itk_process/test_process_load.py @@ -1,14 +1,17 @@ import os + +import cv2 import numpy as np import pytest -import cv2 + from itkit.process.LoadBiomedicalData import ( - LoadImgFromOpenCV, LoadAnnoFromOpenCV, LoadImageFromMHA, - LoadMaskFromMHA + LoadImgFromOpenCV, + LoadMaskFromMHA, ) + @pytest.fixture def temp_opencv_data(tmp_path): img_path = str(tmp_path / "test_img.png") diff --git a/tests/test_models_io.py b/tests/test_models_io.py index f8dfe40..a3580f8 100644 --- a/tests/test_models_io.py +++ b/tests/test_models_io.py @@ -472,8 +472,8 @@ def test_efficientformerv2_io(): @pytest.mark.torch def test_datransunet_io(): """Test DA-TransUNet IO (2D model).""" - from itkit.models.DA_TransUnet.DATransUNet import DA_Transformer from itkit.models.DA_TransUnet.configs import get_r50_b16_config + from itkit.models.DA_TransUnet.DATransUNet import DA_Transformer # Create model config = get_r50_b16_config() @@ -575,7 +575,7 @@ def test_swinumamba_io(): def test_volumevssm_io(): """Test VolumeVSSM IO.""" pytest.importorskip("mamba_ssm", reason="mamba_ssm not installed") - from itkit.models.VMamba.volume_mamba import VolumeVSSM, MambaAggregator1D + from itkit.models.VMamba.volume_mamba import MambaAggregator1D, VolumeVSSM # Mock backbone class MockBackbone(torch.nn.Module): diff --git a/tests/test_onnx_metadata.py b/tests/test_onnx_metadata.py index 6860683..433b7b4 100644 --- a/tests/test_onnx_metadata.py +++ b/tests/test_onnx_metadata.py @@ -1,16 +1,18 @@ import json + import pytest try: - import torch import onnx import onnxruntime + import torch HAS_ORT = True except ImportError: HAS_ORT = False from itkit.mm.inference import ONNXInferBackend + def create_dummy_onnx(path, inference_config_dict=None): # Create a simple model: y = x class DummyModel(torch.nn.Module):