ctkit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
ctkit/dataset.py ADDED
@@ -0,0 +1,1449 @@
1
+ """The :class:`Dataset` class: a cohort of images processed with one protocol.
2
+
3
+ A dataset is lazy by default. It holds paths, not pixels, and loads one image
4
+ at a time, so a cohort that would never fit in memory can still be processed
5
+ in a single call::
6
+
7
+ Dataset.from_directory("data/nifti") \\
8
+ .filter() \\
9
+ .process(config, out_dir="data/processed")
10
+
11
+ Two of the steps in the protocol need statistics pooled over the whole cohort
12
+ — the output shape (a percentile of the observed shapes) and dataset-level
13
+ z-scoring. :meth:`Dataset.process` resolves those automatically, making extra
14
+ passes over the data when it has to.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import fnmatch
20
+ import glob
21
+ import logging
22
+ import os
23
+ from concurrent.futures import ProcessPoolExecutor, as_completed
24
+ from typing import Any, Callable, Iterable, Iterator, Optional, Sequence, Union
25
+
26
+ import numpy as np
27
+
28
+ from .config import ProcessingConfig
29
+ from .image import Number, RadiologyImage, _copy_nifti
30
+ from .io import NIFTI_SUFFIXES, is_dicom_directory
31
+ from .qc import QCCriteria, QCResult, resolve_criteria
32
+ from .validation import (
33
+ CacheMode,
34
+ Integer,
35
+ Interpolator,
36
+ Labels,
37
+ Layout,
38
+ NormalizationMethod,
39
+ OutputFormat,
40
+ OnError,
41
+ PathArg,
42
+ QCLevel,
43
+ QCPreset,
44
+ SliceMode,
45
+ Spacing,
46
+ validate_class,
47
+ )
48
+
49
+ logger = logging.getLogger(__name__)
50
+
51
+ #: Filenames treated as "the image" when scanning a directory of case folders.
52
+ DEFAULT_IMAGE_NAMES = (
53
+ "imaging.nii.gz", "imaging.nii", "image.nii.gz", "image.nii",
54
+ "ct.nii.gz", "ct.nii", "volume.nii.gz",
55
+ )
56
+ #: Filenames treated as "the mask" for a case folder.
57
+ DEFAULT_MASK_NAMES = (
58
+ "segmentation.nii.gz", "segmentation.nii", "mask.nii.gz", "mask.nii",
59
+ "label.nii.gz", "labels.nii.gz", "seg.nii.gz",
60
+ )
61
+ MASK_SUFFIXES = ("_mask", "_seg", "_segmentation", "_label", "_labels")
62
+
63
+
64
+ @validate_class
65
+ class Dataset:
66
+ """A cohort: an ordered collection of :class:`RadiologyImage` objects.
67
+
68
+ Parameters
69
+ ----------
70
+ images:
71
+ Where the cohort comes from. A directory is scanned for the layouts
72
+ :meth:`from_directory` recognizes, a ``.csv``/``.tsv`` is read as a
73
+ metadata table, a glob is expanded, and a list of paths, arrays or
74
+ images becomes exactly those series::
75
+
76
+ Dataset("data/raw")
77
+ Dataset("metadata.csv")
78
+ Dataset("data/raw/*/imaging.nii.gz")
79
+ Dataset(["a.nii.gz", "b.nii.gz"])
80
+
81
+ Use :meth:`from_directory` directly when a layout needs explicit
82
+ `image_pattern`/`mask_pattern` globs.
83
+ name:
84
+ Optional label used in log messages. Defaults to the directory name.
85
+ lazy:
86
+ Keep images unloaded until they are used. Leave this on for cohorts
87
+ that do not fit in memory.
88
+ """
89
+
90
+ def __init__(
91
+ self,
92
+ images: Any,
93
+ name: Optional[str] = None,
94
+ lazy: bool = True,
95
+ **kwargs: Any,
96
+ ) -> None:
97
+ if isinstance(images, (str, os.PathLike)):
98
+ discovered = self._from_source(str(images), name=name, **kwargs)
99
+ images, name = discovered.images, name or discovered.name
100
+ elif isinstance(images, Dataset):
101
+ images = list(images.images)
102
+ elif not isinstance(images, (list, tuple, set, Iterator)):
103
+ # A bare array or NIfTI image is one series, not an iterable of
104
+ # slices, and a string would iterate into characters.
105
+ images = [images]
106
+ elif kwargs:
107
+ images = [
108
+ item if isinstance(item, RadiologyImage) else RadiologyImage(item, **kwargs)
109
+ for item in images
110
+ ]
111
+
112
+ self.name = name
113
+ self.lazy = lazy
114
+ self.images: list = [self._coerce(item, lazy) for item in images]
115
+ self.qc_report = None
116
+ #: Images dropped by the most recent :meth:`filter` call.
117
+ self.rejected: list = []
118
+
119
+ @classmethod
120
+ def _from_source(cls, source: str, name: Optional[str] = None, **kwargs: Any) -> "Dataset":
121
+ """Discover a cohort from a path: a directory, a table, a glob, a file."""
122
+ if os.path.isdir(source):
123
+ if is_dicom_directory(source): # one DICOM series, not a cohort
124
+ return cls([RadiologyImage(source, **kwargs)], name=name)
125
+ return cls.from_directory(source, name=name, **kwargs)
126
+ if source.lower().endswith((".csv", ".tsv")):
127
+ return cls.from_metadata(source, name=name)
128
+ if any(character in source for character in "*?["):
129
+ matches = sorted(glob.glob(source, recursive=True))
130
+ if not matches:
131
+ raise FileNotFoundError(f"No files match {source!r}")
132
+ return cls.from_paths(matches, name=name, **kwargs)
133
+ if not os.path.exists(source):
134
+ raise FileNotFoundError(
135
+ f"No such file or directory: {source}. A Dataset can be built from "
136
+ "a directory of cases, a metadata CSV, a glob, or a list of paths."
137
+ )
138
+ return cls([RadiologyImage(source, **kwargs)], name=name)
139
+
140
+ @staticmethod
141
+ def _coerce(item: Any, lazy: bool) -> RadiologyImage:
142
+ if isinstance(item, RadiologyImage):
143
+ return item
144
+ if isinstance(item, (tuple, list)) and len(item) == 2:
145
+ return RadiologyImage(item[0], mask=item[1], lazy=lazy)
146
+ if isinstance(item, dict):
147
+ return RadiologyImage(**item)
148
+ return RadiologyImage(item, lazy=lazy)
149
+
150
+ # ------------------------------------------------------------------
151
+ # constructors
152
+ # ------------------------------------------------------------------
153
+ @classmethod
154
+ def from_directory(
155
+ cls,
156
+ root: str,
157
+ image_pattern: Optional[str] = None,
158
+ mask_pattern: Optional[str] = None,
159
+ recursive: bool = True,
160
+ name: Optional[str] = None,
161
+ **kwargs: Any,
162
+ ) -> "Dataset":
163
+ """Discover images under `root`.
164
+
165
+ Three layouts are recognized automatically:
166
+
167
+ * one directory per case holding ``imaging.nii.gz`` and (optionally)
168
+ ``segmentation.nii.gz`` — the layout this package writes;
169
+ * a flat directory of NIfTI files, where a file whose name ends in
170
+ ``_mask``/``_seg``/``_label`` is paired with its image;
171
+ * one directory per case holding DICOM slices.
172
+
173
+ `image_pattern` and `mask_pattern` are glob patterns (e.g.
174
+ ``"*_CT.nii.gz"``) that override the automatic rules.
175
+ """
176
+ if not os.path.isdir(root):
177
+ raise NotADirectoryError(f"Not a directory: {root}")
178
+
179
+ if image_pattern or mask_pattern:
180
+ pairs = _scan_with_patterns(root, image_pattern, mask_pattern, recursive)
181
+ else:
182
+ pairs = _scan_case_directories(root)
183
+ if not pairs:
184
+ pairs = _scan_flat_files(root)
185
+ if not pairs:
186
+ pairs = _scan_dicom_directories(root)
187
+
188
+ if not pairs:
189
+ raise FileNotFoundError(
190
+ f"No images found under {root}. Expected per-case directories "
191
+ f"containing one of {DEFAULT_IMAGE_NAMES}, a flat directory of NIfTI "
192
+ "files, or per-case DICOM directories. Pass image_pattern= to "
193
+ "override."
194
+ )
195
+
196
+ images = [
197
+ RadiologyImage(image, mask=mask, series_id=series_id, lazy=True, **kwargs)
198
+ for series_id, image, mask in pairs
199
+ ]
200
+ return cls(images, name=name or os.path.basename(os.path.abspath(root)))
201
+
202
+ @classmethod
203
+ def from_paths(
204
+ cls,
205
+ images: Sequence[Any],
206
+ masks: Optional[Sequence[Any]] = None,
207
+ series_ids: Optional[Sequence[str]] = None,
208
+ name: Optional[str] = None,
209
+ **kwargs: Any,
210
+ ) -> "Dataset":
211
+ """Build from explicit lists of images and (optionally) masks."""
212
+ if masks is not None and len(masks) != len(images):
213
+ raise ValueError(
214
+ f"Got {len(images)} images but {len(masks)} masks; they must match."
215
+ )
216
+ if series_ids is not None and len(series_ids) != len(images):
217
+ raise ValueError(
218
+ f"Got {len(images)} images but {len(series_ids)} series_ids."
219
+ )
220
+
221
+ built = [
222
+ RadiologyImage(
223
+ image,
224
+ mask=None if masks is None else masks[index],
225
+ series_id=None if series_ids is None else series_ids[index],
226
+ lazy=True,
227
+ **kwargs,
228
+ )
229
+ for index, image in enumerate(images)
230
+ ]
231
+ return cls(built, name=name)
232
+
233
+ @classmethod
234
+ def from_metadata(
235
+ cls,
236
+ metadata,
237
+ root: Optional[str] = None,
238
+ image_column: str = "Image",
239
+ mask_column: str = "Mask",
240
+ series_id_column: str = "series_id",
241
+ name: Optional[str] = None,
242
+ ) -> "Dataset":
243
+ """Build from a metadata table (a DataFrame or a path to a CSV).
244
+
245
+ Every column becomes part of each image's :attr:`RadiologyImage.metadata`,
246
+ so acquisition parameters travel with the scan through the pipeline.
247
+ """
248
+ import pandas as pd
249
+
250
+ frame = pd.read_csv(metadata) if isinstance(metadata, str) else metadata
251
+ if image_column not in frame.columns:
252
+ raise KeyError(
253
+ f"Column {image_column!r} not in the metadata "
254
+ f"(columns: {', '.join(map(str, frame.columns))})"
255
+ )
256
+
257
+ def resolve(value: Any) -> Optional[str]:
258
+ if value is None or (isinstance(value, float) and np.isnan(value)):
259
+ return None
260
+ path = str(value)
261
+ if root and not os.path.isabs(path):
262
+ path = os.path.join(root, path)
263
+ return path
264
+
265
+ images = []
266
+ for _, row in frame.iterrows():
267
+ image_path = resolve(row[image_column])
268
+ if image_path is None:
269
+ continue
270
+ mask_path = resolve(row[mask_column]) if mask_column in frame.columns else None
271
+ images.append(
272
+ RadiologyImage(
273
+ image_path,
274
+ mask=mask_path,
275
+ series_id=str(row[series_id_column]) if series_id_column in frame.columns else None,
276
+ metadata=row.to_dict(),
277
+ lazy=True,
278
+ )
279
+ )
280
+ return cls(images, name=name)
281
+
282
+ # ------------------------------------------------------------------
283
+ # sequence protocol
284
+ # ------------------------------------------------------------------
285
+ def __len__(self) -> int:
286
+ return len(self.images)
287
+
288
+ def __iter__(self) -> Iterator[RadiologyImage]:
289
+ return iter(self.images)
290
+
291
+ def __getitem__(self, index):
292
+ if isinstance(index, slice):
293
+ return Dataset(self.images[index], name=self.name, lazy=self.lazy)
294
+ return self.images[index]
295
+
296
+ def __repr__(self) -> str:
297
+ label = f" {self.name!r}" if self.name else ""
298
+ with_masks = sum(1 for image in self.images if image.mask_source or image._mask is not None)
299
+ return f"<Dataset{label} n={len(self)} with_masks={with_masks}>"
300
+
301
+ @property
302
+ def series_ids(self) -> list:
303
+ return [image.series_id for image in self.images]
304
+
305
+ def get(self, series_id: str) -> RadiologyImage:
306
+ """Look up one image by series id."""
307
+ for image in self.images:
308
+ if image.series_id == series_id:
309
+ return image
310
+ raise KeyError(f"No series {series_id!r} in this dataset")
311
+
312
+ # ------------------------------------------------------------------
313
+ # quality control
314
+ # ------------------------------------------------------------------
315
+ def filter(
316
+ self,
317
+ criteria: Union[QCCriteria, QCPreset, None] = None,
318
+ level: QCLevel = "all",
319
+ progress: bool = True,
320
+ keep_report: bool = True,
321
+ **thresholds: Any,
322
+ ) -> "Dataset":
323
+ """Drop series that fail quality control, returning a new dataset.
324
+
325
+ `criteria` is a :class:`QCCriteria`, or the name of a preset
326
+ (``"default"``, ``"radiomics"``, ``"permissive"``). Individual
327
+ thresholds can be given as keywords instead: ``filter(min_slices=25)``.
328
+
329
+ `level` chooses which checks run: ``"metadata"`` reads no pixel data,
330
+ so it is the cheap first pass over a cohort that still has its headers
331
+ and metadata rows; ``"volume"`` checks the reconstructed volumes;
332
+ ``"all"`` (the default) does both.
333
+
334
+ The full pass/fail table, including the measurements taken for series
335
+ that passed, is left on :attr:`qc_report` of the returned dataset (and
336
+ of this one), so exclusions can be reported in a paper. The images that
337
+ failed are kept on :attr:`rejected`, so they can be moved aside or
338
+ deleted afterwards.
339
+ """
340
+ criteria = resolve_criteria(criteria, **thresholds)
341
+ records, kept, rejected = [], [], []
342
+
343
+ for image in _progress(self.images, "Quality control", progress):
344
+ try:
345
+ result = image.check(criteria, level=level)
346
+ except Exception as error: # noqa: BLE001 - unreadable is a failure
347
+ result = QCResult(series_id=image.series_id).fail(f"failed to load: {error}")
348
+ records.append(result.to_dict())
349
+ if result.passed:
350
+ kept.append(image)
351
+ else:
352
+ rejected.append(image)
353
+ logger.info("Excluding %s: %s", image.series_id, result.reason)
354
+ if self.lazy and image.source is not None:
355
+ image.unload()
356
+
357
+ report = _to_frame(records)
358
+ filtered = Dataset(kept, name=self.name, lazy=self.lazy)
359
+ filtered.rejected = rejected
360
+ self.rejected = rejected
361
+ if keep_report:
362
+ self.qc_report = report
363
+ filtered.qc_report = report
364
+
365
+ logger.info(
366
+ "Quality control kept %d of %d series (%d excluded).",
367
+ len(kept), len(self.images), len(self.images) - len(kept),
368
+ )
369
+ return filtered
370
+
371
+ def check(
372
+ self,
373
+ criteria: Union[QCCriteria, QCPreset, None] = None,
374
+ level: QCLevel = "all",
375
+ progress: bool = True,
376
+ **thresholds: Any,
377
+ ):
378
+ """Run quality control over the cohort without dropping anything.
379
+
380
+ Returns the pass/fail table — every series, the outcome, the reason,
381
+ and the measurements behind it — as a DataFrame. Use :meth:`filter`
382
+ when you want the series that passed.
383
+ """
384
+ return self.filter(
385
+ criteria, level=level, progress=progress, **thresholds
386
+ ).qc_report
387
+
388
+ # ------------------------------------------------------------------
389
+ # processing steps, applied to every image
390
+ # ------------------------------------------------------------------
391
+ def orient(self, target: str = "RAS") -> "Dataset":
392
+ """Reorient every image. See :meth:`RadiologyImage.orient`."""
393
+ return self.apply("orient", target=target)
394
+
395
+ def segment(
396
+ self,
397
+ organs: Optional[Sequence[str]] = None,
398
+ task: str = "total",
399
+ tumor_mask: Optional[Any] = None,
400
+ restrict_organs_to_tumor: bool = True,
401
+ fast: bool = False,
402
+ remove_small_blobs: bool = True,
403
+ fill_holes: bool = True,
404
+ morphological_closing: bool = True,
405
+ device: Optional[str] = None,
406
+ output_dir: Optional[PathArg] = None,
407
+ replace: bool = False,
408
+ ) -> "Dataset":
409
+ """Segment every image. See :meth:`RadiologyImage.segment`.
410
+
411
+ `output_dir` keeps the TotalSegmentator files, one subdirectory per
412
+ series, so that the series do not overwrite each other.
413
+ """
414
+ options = dict(
415
+ organs=organs,
416
+ task=task,
417
+ tumor_mask=tumor_mask,
418
+ restrict_organs_to_tumor=restrict_organs_to_tumor,
419
+ fast=fast,
420
+ remove_small_blobs=remove_small_blobs,
421
+ fill_holes=fill_holes,
422
+ morphological_closing=morphological_closing,
423
+ device=device,
424
+ replace=replace,
425
+ )
426
+ if output_dir is None:
427
+ return self.apply("segment", **options)
428
+
429
+ root = os.fspath(output_dir)
430
+
431
+ def segment_one(image: RadiologyImage) -> None:
432
+ image.segment(
433
+ output_dir=os.path.join(root, image.series_id or "image"), **options
434
+ )
435
+
436
+ return self.apply(segment_one)
437
+
438
+ def clip(
439
+ self,
440
+ min_value: Optional[Number] = -200,
441
+ max_value: Optional[Number] = 300,
442
+ ) -> "Dataset":
443
+ """Clamp every image to an intensity window. See :meth:`RadiologyImage.clip`."""
444
+ return self.apply("clip", min_value=min_value, max_value=max_value)
445
+
446
+ def resample(
447
+ self,
448
+ spacing: Spacing = (0.8, 0.8, 3.0),
449
+ interpolator: Interpolator = "linear",
450
+ ) -> "Dataset":
451
+ """Resample every image. See :meth:`RadiologyImage.resample`."""
452
+ return self.apply("resample", spacing=spacing, interpolator=interpolator)
453
+
454
+ def select_slice(
455
+ self,
456
+ mode: SliceMode = "mask",
457
+ index: Optional[Integer] = None,
458
+ label: Optional[Labels] = None,
459
+ keepdims: bool = False,
460
+ ) -> "Dataset":
461
+ """Reduce every volume to one slice. See :meth:`RadiologyImage.select_slice`."""
462
+ return self.apply(
463
+ "select_slice", mode=mode, index=index, label=label, keepdims=keepdims
464
+ )
465
+
466
+ def apply_mask(
467
+ self,
468
+ labels: Optional[Labels] = None,
469
+ crop: bool = True,
470
+ padding: Integer = 5,
471
+ fill_value: Optional[Number] = None,
472
+ ) -> "Dataset":
473
+ """Mask and crop every image. See :meth:`RadiologyImage.apply_mask`."""
474
+ return self.apply(
475
+ "apply_mask", labels=labels, crop=crop, padding=padding, fill_value=fill_value
476
+ )
477
+
478
+ def crop_to_content(
479
+ self, threshold: Optional[Number] = None, padding: Integer = 5
480
+ ) -> "Dataset":
481
+ """Crop every image to its content. See :meth:`RadiologyImage.crop_to_content`."""
482
+ return self.apply("crop_to_content", threshold=threshold, padding=padding)
483
+
484
+ def standardize_size(
485
+ self,
486
+ x: Optional[Integer] = None,
487
+ y: Optional[Integer] = None,
488
+ z: Optional[Integer] = None,
489
+ fill_value: Optional[Number] = None,
490
+ mask_fill_value: Number = 0,
491
+ ) -> "Dataset":
492
+ """Crop/pad every image to a fixed shape. See :meth:`RadiologyImage.standardize_size`.
493
+
494
+ Unlike a protocol, this takes the shape you name; use
495
+ :meth:`shape_percentile` to measure one from the cohort first.
496
+ """
497
+ return self.apply(
498
+ "standardize_size",
499
+ x=x,
500
+ y=y,
501
+ z=z,
502
+ fill_value=fill_value,
503
+ mask_fill_value=mask_fill_value,
504
+ )
505
+
506
+ def normalize(
507
+ self,
508
+ method: NormalizationMethod = "volume",
509
+ mean: Optional[Number] = None,
510
+ std: Optional[Number] = None,
511
+ within_mask: bool = False,
512
+ ) -> "Dataset":
513
+ """Z-score every image. See :meth:`RadiologyImage.normalize`.
514
+
515
+ ``method="dataset"`` needs `mean` and `std` pooled over the cohort;
516
+ :meth:`intensity_statistics` measures them, and :meth:`process`
517
+ resolves them for you.
518
+ """
519
+ return self.apply(
520
+ "normalize", method=method, mean=mean, std=std, within_mask=within_mask
521
+ )
522
+
523
+ # ------------------------------------------------------------------
524
+ # cohort statistics
525
+ # ------------------------------------------------------------------
526
+ def shape_percentile(self, percentile: Number = 95.0, progress: bool = True) -> tuple:
527
+ """Per-axis percentile of the image shapes across the cohort.
528
+
529
+ This is how a common output size is chosen: large enough to contain
530
+ almost every ROI, small enough that the rare huge case is cropped
531
+ rather than every case being padded to it.
532
+ """
533
+ shapes = []
534
+ for image in _progress(self.images, f"Measuring shapes (p{percentile:g})", progress):
535
+ data = image.array
536
+ if data.size and np.all(data == data.flat[0]):
537
+ logger.warning("%s is uniform; excluding it from the shape statistics.",
538
+ image.series_id)
539
+ else:
540
+ shapes.append(image.shape)
541
+ if self.lazy and image.source is not None:
542
+ image.unload()
543
+
544
+ if not shapes:
545
+ raise ValueError("No usable images to measure shapes from")
546
+
547
+ n_axes = max(len(shape) for shape in shapes)
548
+ padded = np.array([list(shape) + [1] * (n_axes - len(shape)) for shape in shapes])
549
+ return tuple(int(np.percentile(padded[:, axis], percentile)) for axis in range(n_axes))
550
+
551
+ def intensity_statistics(self, progress: bool = True) -> dict:
552
+ """Mean and standard deviation pooled over every voxel in the cohort.
553
+
554
+ Computed in one streaming pass with running sums, so it does not need
555
+ the cohort in memory.
556
+ """
557
+ total = 0
558
+ total_sum = 0.0
559
+ total_square_sum = 0.0
560
+
561
+ for image in _progress(self.images, "Pooling intensities", progress):
562
+ data = np.asarray(image.array, dtype=np.float64)
563
+ total += data.size
564
+ total_sum += float(data.sum())
565
+ total_square_sum += float(np.square(data).sum())
566
+ if self.lazy and image.source is not None:
567
+ image.unload()
568
+
569
+ if total == 0:
570
+ raise ValueError("No voxels to compute statistics from")
571
+
572
+ mean = total_sum / total
573
+ variance = max(total_square_sum / total - mean * mean, 0.0)
574
+ return {"mean": mean, "std": float(np.sqrt(variance)), "n_voxels": total}
575
+
576
+ def statistics(self, progress: bool = True):
577
+ """Per-image summary statistics as a DataFrame."""
578
+ records = []
579
+ for image in _progress(self.images, "Summarizing", progress):
580
+ records.append(image.statistics())
581
+ if self.lazy and image.source is not None:
582
+ image.unload()
583
+ return _to_frame(records)
584
+
585
+ def check_consistency(
586
+ self,
587
+ mean_tolerance: float = 100.0,
588
+ std_fraction_tolerance: float = 0.15,
589
+ progress: bool = True,
590
+ ) -> list:
591
+ """Flag images whose intensity distribution is unlike the rest.
592
+
593
+ A series that survives every other check can still be wrong — a
594
+ different contrast phase, a failed rescale, an inverted intensity
595
+ scale. Comparing each volume against the cohort median catches those.
596
+ """
597
+ frame = self.statistics(progress=progress)
598
+ if frame is None or len(frame) == 0:
599
+ return []
600
+
601
+ median_mean = float(np.median(frame["mean"]))
602
+ median_std = float(np.median(frame["std"]))
603
+ problems = []
604
+ for _, row in frame.iterrows():
605
+ if abs(row["mean"] - median_mean) > mean_tolerance:
606
+ problems.append(
607
+ f"{row['series_id']}: mean {row['mean']:.1f} differs from the cohort "
608
+ f"median {median_mean:.1f}"
609
+ )
610
+ if abs(row["std"] - median_std) > std_fraction_tolerance * median_std:
611
+ problems.append(
612
+ f"{row['series_id']}: std {row['std']:.1f} differs from the cohort "
613
+ f"median {median_std:.1f}"
614
+ )
615
+
616
+ if problems:
617
+ logger.warning("Intensity consistency check found %d issues:", len(problems))
618
+ for problem in problems:
619
+ logger.warning(" %s", problem)
620
+ else:
621
+ logger.info("Intensity consistency check passed for all %d series.", len(frame))
622
+ return problems
623
+
624
+ # ------------------------------------------------------------------
625
+ # processing
626
+ # ------------------------------------------------------------------
627
+ def resolve_config(
628
+ self, config: ProcessingConfig, progress: bool = True
629
+ ) -> ProcessingConfig:
630
+ """Fill in the config fields that depend on the whole cohort.
631
+
632
+ Resolves ``target_shape`` (from the shape percentile) and
633
+ ``dataset_mean``/``dataset_std``, each measured at the point in the
634
+ pipeline where the corresponding step runs.
635
+ """
636
+ resolved = config
637
+
638
+ if config.standardize_size and config.target_shape is None:
639
+ logger.info(
640
+ "target_shape is unset: measuring the p%g shape across %d series.",
641
+ config.shape_percentile, len(self),
642
+ )
643
+ probe = resolved.replace(standardize_size=False, normalize=False)
644
+ shapes = self._processed_view(probe, progress=progress).shape_percentile(
645
+ config.shape_percentile, progress=progress
646
+ )
647
+ shapes = (list(shapes) + [None, None, None])[:3]
648
+ resolved = resolved.replace(target_shape=tuple(shapes))
649
+ logger.info("Resolved target_shape=%s", tuple(shapes))
650
+
651
+ if (config.normalize and config.normalization_method == "dataset"
652
+ and config.dataset_mean is None):
653
+ logger.info("Pooling cohort intensity statistics for dataset normalization.")
654
+ probe = resolved.replace(normalize=False)
655
+ stats = self._processed_view(probe, progress=progress).intensity_statistics(
656
+ progress=progress
657
+ )
658
+ resolved = resolved.replace(
659
+ dataset_mean=stats["mean"], dataset_std=stats["std"]
660
+ )
661
+ logger.info(
662
+ "Resolved dataset_mean=%.4f dataset_std=%.4f", stats["mean"], stats["std"]
663
+ )
664
+
665
+ return resolved
666
+
667
+ def _processed_view(self, config: ProcessingConfig, progress: bool = True) -> "Dataset":
668
+ """A dataset of images processed with `config`, materialized on demand.
669
+
670
+ Used for the measurement passes. The processed volumes are produced one
671
+ at a time and released, never written to disk.
672
+ """
673
+ return _ProcessedView(self, config, progress=progress)
674
+
675
+ @staticmethod
676
+ def _split_config(config: ProcessingConfig):
677
+ """Split a protocol into the part before the cohort-level steps and the rest.
678
+
679
+ The expensive steps — segmentation above all — live in the first half.
680
+ Running them once, caching the result, and applying only the second half
681
+ afterwards keeps a cohort-level protocol to a single segmentation pass.
682
+ """
683
+ head = config.replace(standardize_size=False, normalize=False)
684
+ tail = config.replace(
685
+ orient=False,
686
+ segment=False,
687
+ clip=False,
688
+ crop_to_content=False,
689
+ resample=False,
690
+ mask=False,
691
+ dimensionality="3D", # the slice, if any, was already selected
692
+ )
693
+ return head, tail
694
+
695
+ def process(
696
+ self,
697
+ config: Optional[Any] = None,
698
+ out_dir: Optional[PathArg] = None,
699
+ workers: Integer = 1,
700
+ progress: bool = True,
701
+ layout: Layout = "case_dirs",
702
+ skip_existing: bool = False,
703
+ on_error: OnError = "warn",
704
+ cache: CacheMode = "auto",
705
+ cache_dir: Optional[PathArg] = None,
706
+ **overrides: Any,
707
+ ) -> "Dataset":
708
+ """Run the protocol over the cohort.
709
+
710
+ Parameters
711
+ ----------
712
+ config:
713
+ The protocol: a :class:`ProcessingConfig`, the name of a collection
714
+ whose curated protocol to use, or a path to a saved YAML protocol.
715
+ Defaults to :class:`ProcessingConfig` defaults.
716
+ out_dir:
717
+ Where to write the processed images. Only the final image (and its
718
+ mask) is written — no intermediate files. When omitted, results
719
+ are kept in memory and returned, which needs the whole cohort to
720
+ fit in RAM.
721
+ workers:
722
+ Processes to run in parallel. Leave at 1 when segmenting on a
723
+ single GPU.
724
+ layout:
725
+ ``"case_dirs"`` writes ``out_dir/<series_id>/imaging.nii.gz``;
726
+ ``"flat"`` writes ``out_dir/<series_id>.nii.gz``.
727
+ skip_existing:
728
+ Skip series whose output already exists, so an interrupted run can
729
+ be resumed.
730
+ on_error:
731
+ ``"warn"`` logs and continues to the next series, ``"raise"`` stops.
732
+ cache:
733
+ How to handle protocols with cohort-level steps, which need more
734
+ than one look at the data. ``"auto"`` (the default) stages the
735
+ expensive part of the pipeline in a temporary directory when the
736
+ protocol segments, so TotalSegmentator runs once instead of three
737
+ times; ``"disk"`` always stages; ``"none"`` recomputes instead,
738
+ trading time for temporary disk. The staging directory is deleted
739
+ before this returns either way.
740
+ cache_dir:
741
+ Where to stage. Defaults to the system temporary directory.
742
+
743
+ Returns
744
+ -------
745
+ Dataset
746
+ The processed cohort. Backed by the written files when `out_dir`
747
+ was given.
748
+ """
749
+ config = ProcessingConfig.resolve(config, **overrides)
750
+ config.validate()
751
+
752
+ if config.needs_dataset_pass:
753
+ logger.info(
754
+ "This protocol has cohort-level steps (%s), so the data is read more "
755
+ "than once. Set target_shape/dataset_mean/dataset_std in the config to "
756
+ "process in a single pass.",
757
+ ", ".join(
758
+ step for step in ("standardize_size", "normalize")
759
+ if step in config.steps
760
+ ),
761
+ )
762
+ staging = cache == "disk" or (cache == "auto" and config.segment)
763
+ if staging:
764
+ return self._process_staged(
765
+ config, out_dir, workers, progress, layout, skip_existing,
766
+ on_error, cache_dir,
767
+ )
768
+ config = self.resolve_config(config, progress=progress)
769
+
770
+ if out_dir:
771
+ os.makedirs(out_dir, exist_ok=True)
772
+ config.to_yaml(os.path.join(out_dir, "processing_config.yaml"))
773
+
774
+ if workers > 1 and out_dir:
775
+ results = self._process_parallel(
776
+ config, out_dir, workers, progress, layout, skip_existing, on_error
777
+ )
778
+ else:
779
+ if workers > 1:
780
+ logger.warning(
781
+ "workers>1 needs out_dir (results are passed back as files); "
782
+ "processing sequentially instead."
783
+ )
784
+ results = self._process_sequential(
785
+ config, out_dir, progress, layout, skip_existing, on_error
786
+ )
787
+
788
+ processed = Dataset(
789
+ [item for item in results if item is not None],
790
+ name=self.name,
791
+ lazy=bool(out_dir),
792
+ )
793
+ if out_dir:
794
+ self._write_manifest(processed, out_dir)
795
+ return processed
796
+
797
+ def _process_staged(
798
+ self, config, out_dir, workers, progress, layout, skip_existing,
799
+ on_error, cache_dir,
800
+ ) -> "Dataset":
801
+ """Run a cohort-level protocol with one pass over the expensive steps.
802
+
803
+ The pipeline is split in two. The first half — orientation,
804
+ segmentation, clipping, resampling, masking — runs once into a
805
+ temporary directory. The cohort statistics are then measured from that
806
+ staged output, and only the cheap second half (crop/pad and z-scoring)
807
+ is applied to produce the final files. The staging directory is removed
808
+ before returning.
809
+ """
810
+ import shutil
811
+ import tempfile
812
+
813
+ head, tail = self._split_config(config)
814
+ staging = tempfile.mkdtemp(prefix="ctkit_stage_", dir=cache_dir)
815
+ logger.info(
816
+ "Staging the per-image steps (%s) once in a temporary directory, so they "
817
+ "are not repeated for each cohort-level measurement.",
818
+ ", ".join(head.steps) or "none",
819
+ )
820
+ try:
821
+ staged = self.process(
822
+ head,
823
+ out_dir=staging,
824
+ workers=workers,
825
+ progress=progress,
826
+ layout="case_dirs",
827
+ on_error=on_error,
828
+ cache="none",
829
+ )
830
+ # Measure on the staged images. `tail` has the per-image steps
831
+ # switched off, so the probe passes over the staged data are
832
+ # no-ops and nothing expensive is repeated.
833
+ resolved_tail = staged.resolve_config(tail, progress=progress)
834
+ resolved = config.replace(
835
+ target_shape=resolved_tail.target_shape,
836
+ dataset_mean=resolved_tail.dataset_mean,
837
+ dataset_std=resolved_tail.dataset_std,
838
+ )
839
+ processed = staged.process(
840
+ resolved_tail,
841
+ out_dir=out_dir,
842
+ workers=workers,
843
+ progress=progress,
844
+ layout=layout,
845
+ skip_existing=skip_existing,
846
+ on_error=on_error,
847
+ cache="none",
848
+ )
849
+ if out_dir:
850
+ # Record the protocol as a whole, not just the half applied last.
851
+ resolved.to_yaml(os.path.join(out_dir, "processing_config.yaml"))
852
+ return processed
853
+ finally:
854
+ shutil.rmtree(staging, ignore_errors=True)
855
+
856
+ def _process_sequential(
857
+ self, config, out_dir, progress, layout, skip_existing, on_error
858
+ ) -> list:
859
+ results = []
860
+ for image in _progress(self.images, "Processing", progress):
861
+ try:
862
+ results.append(
863
+ _process_one_image(image, config, out_dir, layout, skip_existing)
864
+ )
865
+ except Exception as error: # noqa: BLE001 - one bad series must not stop a cohort
866
+ if on_error == "raise":
867
+ raise
868
+ logger.error("Failed to process %s: %s", image.series_id, error)
869
+ results.append(None)
870
+ finally:
871
+ if self.lazy and image.source is not None and out_dir:
872
+ image.unload()
873
+ return results
874
+
875
+ def _process_parallel(
876
+ self, config, out_dir, workers, progress, layout, skip_existing, on_error
877
+ ) -> list:
878
+ # Workers are handed a path and rebuild the image from it, so anything
879
+ # that is not fully described by its path has to stay in this process.
880
+ unpicklable = [
881
+ image
882
+ for image in self.images
883
+ if image.source is None or isinstance(image, _DeferredProxy)
884
+ ]
885
+ if unpicklable:
886
+ logger.warning(
887
+ "%d images hold in-memory data or a deferred step and cannot be "
888
+ "sent to worker processes; processing sequentially instead.",
889
+ len(unpicklable),
890
+ )
891
+ return self._process_sequential(
892
+ config, out_dir, progress, layout, skip_existing, on_error
893
+ )
894
+
895
+ payloads = [
896
+ {
897
+ "source": str(image.source),
898
+ "mask_source": None if image.mask_source is None else str(image.mask_source),
899
+ "series_id": image.series_id,
900
+ "metadata": image.metadata,
901
+ "masked_roi": image._masked_roi,
902
+ "config": config.to_dict(),
903
+ "out_dir": out_dir,
904
+ "layout": layout,
905
+ "skip_existing": skip_existing,
906
+ }
907
+ for image in self.images
908
+ ]
909
+
910
+ results: list = [None] * len(payloads)
911
+ with ProcessPoolExecutor(max_workers=workers) as pool:
912
+ futures = {
913
+ pool.submit(_process_payload, payload): index
914
+ for index, payload in enumerate(payloads)
915
+ }
916
+ for future in _progress(
917
+ as_completed(futures), "Processing", progress, total=len(futures)
918
+ ):
919
+ index = futures[future]
920
+ try:
921
+ outcome = future.result()
922
+ except Exception as error: # noqa: BLE001
923
+ if on_error == "raise":
924
+ raise
925
+ logger.error(
926
+ "Failed to process %s: %s", payloads[index]["series_id"], error
927
+ )
928
+ continue
929
+ results[index] = RadiologyImage(
930
+ outcome["image_path"],
931
+ mask=outcome.get("mask_path"),
932
+ series_id=outcome["series_id"],
933
+ metadata=outcome.get("metadata") or {},
934
+ lazy=True,
935
+ )
936
+ return results
937
+
938
+ def _write_manifest(self, processed: "Dataset", out_dir: str) -> None:
939
+ records = []
940
+ for image in processed:
941
+ record = {"series_id": image.series_id, "image": image.source}
942
+ if image.mask_source:
943
+ record["mask"] = image.mask_source
944
+ record.update(
945
+ {key: value for key, value in image.metadata.items()
946
+ if not isinstance(value, (list, dict, np.ndarray))}
947
+ )
948
+ records.append(record)
949
+ frame = _to_frame(records)
950
+ if frame is not None:
951
+ frame.to_csv(os.path.join(out_dir, "manifest.csv"), index=False)
952
+
953
+ # ------------------------------------------------------------------
954
+ # feature extraction
955
+ # ------------------------------------------------------------------
956
+ def radiomics(
957
+ self,
958
+ labels: Labels = (1, 2),
959
+ params: Optional[Union[str, dict]] = None,
960
+ out_csv: Optional[PathArg] = None,
961
+ progress: bool = True,
962
+ on_error: OnError = "warn",
963
+ ):
964
+ """Extract radiomic features for every series, as a DataFrame."""
965
+ from .features import extract_features
966
+
967
+ records = []
968
+ for image in _progress(self.images, "Extracting features", progress):
969
+ try:
970
+ features = extract_features(image, labels=labels, params=params)
971
+ except Exception as error: # noqa: BLE001
972
+ if on_error == "raise":
973
+ raise
974
+ logger.error("Feature extraction failed for %s: %s", image.series_id, error)
975
+ continue
976
+ finally:
977
+ if self.lazy and image.source is not None:
978
+ image.unload()
979
+ records.append(features)
980
+
981
+ frame = _to_frame(records)
982
+ if out_csv and frame is not None:
983
+ os.makedirs(os.path.dirname(os.path.abspath(out_csv)) or ".", exist_ok=True)
984
+ frame.to_csv(out_csv, index=False)
985
+ logger.info("Wrote %d feature rows to %s", len(frame), out_csv)
986
+ return frame
987
+
988
+ # ------------------------------------------------------------------
989
+ # misc
990
+ # ------------------------------------------------------------------
991
+ def apply(self, step: Union[str, Callable], *args: Any, **kwargs: Any) -> "Dataset":
992
+ """Apply one processing step to every image, returning a new dataset.
993
+
994
+ `step` is the name of a :class:`RadiologyImage` method (``"orient"``,
995
+ ``"clip"``, ...) or a callable taking an image.
996
+
997
+ For images backed by a path the step is *deferred*: it is recorded now
998
+ and run as each image is read, so a cohort that does not fit in memory
999
+ still works and repeated calls compose into a single pass. Images that
1000
+ are already in memory have nowhere to be re-read from, so the step is
1001
+ applied to them immediately, in place.
1002
+ """
1003
+ if not callable(step) and not callable(getattr(RadiologyImage, str(step), None)):
1004
+ raise AttributeError(
1005
+ f"{step!r} is not a RadiologyImage method. Available steps: "
1006
+ "orient, segment, clip, resample, select_slice, apply_mask, "
1007
+ "crop_to_content, standardize_size, normalize."
1008
+ )
1009
+
1010
+ def run(image: RadiologyImage) -> None:
1011
+ if callable(step):
1012
+ step(image)
1013
+ else:
1014
+ getattr(image, step)(*args, **kwargs)
1015
+
1016
+ images = [
1017
+ _DeferredProxy(image, run) if image.source is not None else _in_place(image, run)
1018
+ for image in self.images
1019
+ ]
1020
+ applied = Dataset(images, name=self.name, lazy=self.lazy)
1021
+ applied.qc_report = self.qc_report
1022
+ return applied
1023
+
1024
+ def map(
1025
+ self, function: Callable[[RadiologyImage], Any], progress: bool = True
1026
+ ) -> list:
1027
+ """Apply `function` to each image, returning the results."""
1028
+ outputs = []
1029
+ for image in _progress(self.images, "Mapping", progress):
1030
+ outputs.append(function(image))
1031
+ if self.lazy and image.source is not None:
1032
+ image.unload()
1033
+ return outputs
1034
+
1035
+ def save(
1036
+ self,
1037
+ out_dir: PathArg,
1038
+ layout: Layout = "case_dirs",
1039
+ output_format: OutputFormat = "nifti",
1040
+ compress: bool = True,
1041
+ progress: bool = True,
1042
+ ) -> "Dataset":
1043
+ """Write every image in the dataset to `out_dir`."""
1044
+ os.makedirs(out_dir, exist_ok=True)
1045
+ written = []
1046
+ for image in _progress(self.images, "Saving", progress):
1047
+ path = _output_path(image, out_dir, layout, output_format, compress)
1048
+ image.save(path, output_format=output_format, compress=compress)
1049
+ written.append(image)
1050
+ return Dataset(written, name=self.name, lazy=self.lazy)
1051
+
1052
+ @property
1053
+ def metadata(self):
1054
+ """Per-image metadata as a DataFrame."""
1055
+ return _to_frame(
1056
+ [{"series_id": image.series_id, **image.metadata} for image in self.images]
1057
+ )
1058
+
1059
+
1060
+ class _ProcessedView(Dataset):
1061
+ """A dataset whose images are processed with a config as they are read.
1062
+
1063
+ Used for the measurement passes in :meth:`Dataset.resolve_config`. Nothing
1064
+ is cached: each access reprocesses, which trades compute for memory and
1065
+ keeps the promise that no intermediate files are written.
1066
+ """
1067
+
1068
+ def __init__(self, source: Dataset, config: ProcessingConfig, progress: bool = True):
1069
+ self._source = source
1070
+ self._config = config
1071
+ self.name = source.name
1072
+ self.lazy = True
1073
+ self.qc_report = None
1074
+ self.rejected = []
1075
+ self.images = [_ProcessedProxy(image, config) for image in source.images]
1076
+
1077
+
1078
+ class _DeferredProxy(RadiologyImage):
1079
+ """A :class:`RadiologyImage` that runs `apply` the first time it loads.
1080
+
1081
+ This is what lets a step be recorded against a cohort without reading it:
1082
+ the work happens when the image is next needed, and unloading throws it
1083
+ away again, so only one image is ever in memory.
1084
+ """
1085
+
1086
+ def __init__(self, source_image: RadiologyImage, apply: Callable[[RadiologyImage], Any]):
1087
+ self._source_image = source_image
1088
+ self._apply = apply
1089
+ super().__init__(
1090
+ source_image.source if source_image.source is not None else np.zeros((1, 1, 1)),
1091
+ mask=source_image.mask_source,
1092
+ series_id=source_image.series_id,
1093
+ metadata=dict(source_image.metadata),
1094
+ masked_roi=source_image._masked_roi,
1095
+ lazy=True,
1096
+ )
1097
+ # An in-memory source has no path to lazily read back from, so mark it
1098
+ # unloaded here and copy the pixels in load().
1099
+ self._image = None
1100
+ self._mask = None
1101
+ self._loaded = False
1102
+
1103
+ def load(self) -> "RadiologyImage":
1104
+ if not self._loaded:
1105
+ source = self._source_image
1106
+ if source.source is not None and not isinstance(source, _DeferredProxy):
1107
+ super().load()
1108
+ else:
1109
+ # Read through the source itself, so a step it is deferring in
1110
+ # turn runs before this one. Its pixels are copied rather than
1111
+ # shared, so the chain does not mutate what it reads from.
1112
+ was_loaded = source._loaded
1113
+ source.load()
1114
+ self._image = _copy_nifti(source._image)
1115
+ self._mask = _copy_nifti(source._mask)
1116
+ self._loaded = True
1117
+ self.metadata.update(source.metadata)
1118
+ # Carry the provenance of the steps that ran before this one.
1119
+ self.history = list(source.history) + self.history
1120
+ if not was_loaded and source.source is not None:
1121
+ source.unload()
1122
+ self._apply(self)
1123
+ return self
1124
+
1125
+ def unload(self) -> "RadiologyImage":
1126
+ self._image = None
1127
+ self._mask = None
1128
+ self._loaded = False
1129
+ self.history = []
1130
+ return self
1131
+
1132
+
1133
+ class _ProcessedProxy(_DeferredProxy):
1134
+ """A :class:`RadiologyImage` that runs a config the first time it loads."""
1135
+
1136
+ def __init__(self, source_image: RadiologyImage, config: ProcessingConfig):
1137
+ self._config = config
1138
+ super().__init__(source_image, lambda image: image.process(config))
1139
+
1140
+
1141
+ def _in_place(image: RadiologyImage, run: Callable[[RadiologyImage], Any]) -> RadiologyImage:
1142
+ """Run a step now, for an image that cannot be re-read from disk."""
1143
+ run(image)
1144
+ return image
1145
+
1146
+
1147
+ # ----------------------------------------------------------------------
1148
+ # worker entry points
1149
+ # ----------------------------------------------------------------------
1150
+ def _process_one_image(
1151
+ image: RadiologyImage,
1152
+ config: ProcessingConfig,
1153
+ out_dir: Optional[str],
1154
+ layout: str,
1155
+ skip_existing: bool,
1156
+ ) -> RadiologyImage:
1157
+ if out_dir is None:
1158
+ return image.copy().process(config)
1159
+
1160
+ path = _output_path(image, out_dir, layout, config.output_format, config.compress)
1161
+ mask_path = _mask_output_path(path, layout)
1162
+
1163
+ if skip_existing and os.path.exists(path):
1164
+ logger.debug("Skipping %s: %s already exists", image.series_id, path)
1165
+ return RadiologyImage(
1166
+ path,
1167
+ mask=mask_path if os.path.exists(mask_path) else None,
1168
+ series_id=image.series_id,
1169
+ metadata=image.metadata,
1170
+ lazy=True,
1171
+ )
1172
+
1173
+ processed = image.copy().process(config)
1174
+ written = processed.save(
1175
+ path,
1176
+ mask_path=mask_path,
1177
+ output_format=config.output_format,
1178
+ compress=config.compress,
1179
+ save_mask=config.save_mask,
1180
+ )
1181
+ return RadiologyImage(
1182
+ written,
1183
+ mask=mask_path if os.path.exists(mask_path) else None,
1184
+ series_id=processed.series_id,
1185
+ metadata=processed.metadata,
1186
+ lazy=True,
1187
+ )
1188
+
1189
+
1190
+ def _process_payload(payload: dict) -> dict:
1191
+ """Module-level worker so :class:`ProcessPoolExecutor` can pickle it."""
1192
+ config = ProcessingConfig.from_dict(payload["config"])
1193
+ image = RadiologyImage(
1194
+ payload["source"],
1195
+ mask=payload["mask_source"],
1196
+ series_id=payload["series_id"],
1197
+ metadata=payload["metadata"],
1198
+ masked_roi=payload["masked_roi"],
1199
+ lazy=True,
1200
+ )
1201
+ processed = _process_one_image(
1202
+ image, config, payload["out_dir"], payload["layout"], payload["skip_existing"]
1203
+ )
1204
+ return {
1205
+ "series_id": processed.series_id,
1206
+ "image_path": processed.source,
1207
+ "mask_path": processed.mask_source,
1208
+ "metadata": {
1209
+ key: value for key, value in processed.metadata.items()
1210
+ if isinstance(value, (str, int, float, bool, type(None)))
1211
+ },
1212
+ }
1213
+
1214
+
1215
+ # ----------------------------------------------------------------------
1216
+ # disposing of rejected series
1217
+ # ----------------------------------------------------------------------
1218
+ def files_of(image: RadiologyImage) -> list:
1219
+ """The files on disk that belong to `image`, as paths.
1220
+
1221
+ A case directory named after the series (the layout this package writes,
1222
+ and the one TCIA conversions produce) counts as one unit, so its mask and
1223
+ any sidecar files travel with the image. Images that were never read from
1224
+ disk have no files.
1225
+ """
1226
+ if image.source is None:
1227
+ return []
1228
+
1229
+ source = os.path.abspath(str(image.source))
1230
+ parent = os.path.dirname(source)
1231
+ if os.path.isdir(source):
1232
+ return [source]
1233
+ if image.series_id and os.path.basename(parent) == image.series_id:
1234
+ return [parent]
1235
+
1236
+ paths = [source]
1237
+ if image.mask_source is not None:
1238
+ paths.append(os.path.abspath(str(image.mask_source)))
1239
+ return paths
1240
+
1241
+
1242
+ def discard(
1243
+ images: Iterable[RadiologyImage],
1244
+ destination: Optional[str] = None,
1245
+ delete: bool = False,
1246
+ ) -> list:
1247
+ """Move the files of `images` to `destination`, or delete them.
1248
+
1249
+ This is the destructive half of quality control: once a series is excluded,
1250
+ its files are usually in the way. Moving is the reversible option; `delete`
1251
+ has to be asked for explicitly. Returns the paths acted on.
1252
+ """
1253
+ import shutil
1254
+
1255
+ if delete == bool(destination):
1256
+ raise ValueError("discard() takes either a destination or delete=True")
1257
+
1258
+ if destination:
1259
+ os.makedirs(destination, exist_ok=True)
1260
+
1261
+ acted: list = []
1262
+ for image in images:
1263
+ for path in files_of(image):
1264
+ if not os.path.exists(path):
1265
+ continue
1266
+ if delete:
1267
+ if os.path.isdir(path):
1268
+ shutil.rmtree(path, ignore_errors=True)
1269
+ else:
1270
+ os.remove(path)
1271
+ else:
1272
+ target = os.path.join(destination, os.path.basename(path))
1273
+ if os.path.exists(target):
1274
+ logger.warning("%s already exists; leaving %s in place.", target, path)
1275
+ continue
1276
+ shutil.move(path, target)
1277
+ acted.append(path)
1278
+
1279
+ logger.info(
1280
+ "%s %d paths from the series that failed quality control.",
1281
+ "Deleted" if delete else f"Moved to {destination}:", len(acted),
1282
+ )
1283
+ return acted
1284
+
1285
+
1286
+ # ----------------------------------------------------------------------
1287
+ # path helpers
1288
+ # ----------------------------------------------------------------------
1289
+ def _output_path(
1290
+ image: RadiologyImage,
1291
+ out_dir: str,
1292
+ layout: str,
1293
+ output_format: str,
1294
+ compress: bool,
1295
+ ) -> str:
1296
+ from .io import resolve_output_path
1297
+
1298
+ series_id = image.series_id or "image"
1299
+ if layout == "case_dirs":
1300
+ base = os.path.join(out_dir, series_id, "imaging")
1301
+ else:
1302
+ base = os.path.join(out_dir, series_id)
1303
+ return resolve_output_path(base, output_format=output_format, compress=compress)
1304
+
1305
+
1306
+ def _mask_output_path(image_path: str, layout: str = "flat") -> str:
1307
+ """Where the mask goes for a given image path.
1308
+
1309
+ ``case_dirs`` writes ``<case>/segmentation.nii.gz`` next to
1310
+ ``<case>/imaging.nii.gz``, which is the layout
1311
+ :meth:`Dataset.from_directory` reads back; ``flat`` appends ``_mask``.
1312
+ """
1313
+ for extension in (".nii.gz", ".nii", ".npy"):
1314
+ if image_path.lower().endswith(extension):
1315
+ stem = image_path[: -len(extension)]
1316
+ if layout == "case_dirs":
1317
+ return os.path.join(os.path.dirname(stem), "segmentation") + extension
1318
+ return stem + "_mask" + extension
1319
+ return image_path + "_mask"
1320
+
1321
+
1322
+ def _scan_case_directories(root: str) -> list:
1323
+ """``root/<series_id>/imaging.nii.gz`` (+ optional mask)."""
1324
+ pairs = []
1325
+ for entry in sorted(os.listdir(root)):
1326
+ case_dir = os.path.join(root, entry)
1327
+ if not os.path.isdir(case_dir):
1328
+ continue
1329
+ names = set(os.listdir(case_dir))
1330
+ image_name = next((name for name in DEFAULT_IMAGE_NAMES if name in names), None)
1331
+ if image_name is None:
1332
+ continue
1333
+ mask_name = next((name for name in DEFAULT_MASK_NAMES if name in names), None)
1334
+ if mask_name is None:
1335
+ # Fall back to any file whose stem marks it as a mask, e.g.
1336
+ # "imaging_mask.nii.gz" from the flat output layout.
1337
+ from .io import _strip_image_suffix
1338
+
1339
+ mask_name = next(
1340
+ (
1341
+ name for name in sorted(names)
1342
+ if name.lower().endswith(NIFTI_SUFFIXES)
1343
+ and _strip_image_suffix(name).lower().endswith(MASK_SUFFIXES)
1344
+ ),
1345
+ None,
1346
+ )
1347
+ pairs.append((
1348
+ entry,
1349
+ os.path.join(case_dir, image_name),
1350
+ os.path.join(case_dir, mask_name) if mask_name else None,
1351
+ ))
1352
+ return pairs
1353
+
1354
+
1355
+ def _scan_flat_files(root: str) -> list:
1356
+ """A directory of NIfTI files, pairing ``x.nii.gz`` with ``x_mask.nii.gz``."""
1357
+ from .io import _strip_image_suffix
1358
+
1359
+ files = sorted(
1360
+ name for name in os.listdir(root)
1361
+ if name.lower().endswith(NIFTI_SUFFIXES) and os.path.isfile(os.path.join(root, name))
1362
+ )
1363
+ stems = {_strip_image_suffix(name): name for name in files}
1364
+
1365
+ pairs = []
1366
+ for stem, name in stems.items():
1367
+ if any(stem.lower().endswith(suffix) for suffix in MASK_SUFFIXES):
1368
+ continue
1369
+ mask_name = next(
1370
+ (stems[stem + suffix] for suffix in MASK_SUFFIXES if stem + suffix in stems),
1371
+ None,
1372
+ )
1373
+ pairs.append((
1374
+ stem,
1375
+ os.path.join(root, name),
1376
+ os.path.join(root, mask_name) if mask_name else None,
1377
+ ))
1378
+ return pairs
1379
+
1380
+
1381
+ def _scan_dicom_directories(root: str) -> list:
1382
+ """``root/<series_id>/*.dcm``."""
1383
+ pairs = []
1384
+ for entry in sorted(os.listdir(root)):
1385
+ case_dir = os.path.join(root, entry)
1386
+ if not os.path.isdir(case_dir):
1387
+ continue
1388
+ has_dicom = any(
1389
+ name.lower().endswith((".dcm", ".ima"))
1390
+ for _, _, names in os.walk(case_dir)
1391
+ for name in names
1392
+ )
1393
+ if has_dicom:
1394
+ pairs.append((entry, case_dir, None))
1395
+ return pairs
1396
+
1397
+
1398
+ def _scan_with_patterns(
1399
+ root: str,
1400
+ image_pattern: Optional[str],
1401
+ mask_pattern: Optional[str],
1402
+ recursive: bool,
1403
+ ) -> list:
1404
+ from .io import _strip_image_suffix
1405
+
1406
+ matches, mask_matches = [], []
1407
+ walker = os.walk(root) if recursive else [(root, [], os.listdir(root))]
1408
+ for directory, _, names in walker:
1409
+ for name in sorted(names):
1410
+ path = os.path.join(directory, name)
1411
+ if image_pattern and fnmatch.fnmatch(name, image_pattern):
1412
+ matches.append(path)
1413
+ elif mask_pattern and fnmatch.fnmatch(name, mask_pattern):
1414
+ mask_matches.append(path)
1415
+
1416
+ by_directory = {}
1417
+ for path in mask_matches:
1418
+ by_directory.setdefault(os.path.dirname(path), []).append(path)
1419
+
1420
+ pairs = []
1421
+ for path in sorted(matches):
1422
+ directory = os.path.dirname(path)
1423
+ series_id = (
1424
+ os.path.basename(directory)
1425
+ if os.path.abspath(directory) != os.path.abspath(root)
1426
+ else _strip_image_suffix(os.path.basename(path))
1427
+ )
1428
+ candidates = by_directory.get(directory) or []
1429
+ pairs.append((series_id, path, candidates[0] if candidates else None))
1430
+ return pairs
1431
+
1432
+
1433
+ # ----------------------------------------------------------------------
1434
+ # small utilities
1435
+ # ----------------------------------------------------------------------
1436
+ def _progress(iterable, description: str, enabled: bool, total: Optional[int] = None):
1437
+ if not enabled:
1438
+ return iterable
1439
+ try:
1440
+ from tqdm.auto import tqdm
1441
+ except ImportError: # pragma: no cover - tqdm is a dependency, but be safe
1442
+ return iterable
1443
+ return tqdm(iterable, desc=description, total=total)
1444
+
1445
+
1446
+ def _to_frame(records: list):
1447
+ import pandas as pd
1448
+
1449
+ return pd.DataFrame(records)