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/config.py ADDED
@@ -0,0 +1,453 @@
1
+ """Configuration objects for the processing pipeline.
2
+
3
+ A :class:`ProcessingConfig` is a plain, serializable description of *what* the
4
+ pipeline does. It carries no data and no state, so it can be written to YAML,
5
+ committed alongside a paper, and handed to a collaborator who wants to
6
+ reproduce a dataset exactly.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import os
13
+ from dataclasses import asdict, dataclass, fields
14
+ from typing import Any, Optional, Sequence, Union
15
+
16
+ from .constants import tcia_dataset_to_info
17
+ from .validation import (
18
+ OUTPUT_FORMATS,
19
+ Dimensionality,
20
+ NormalizationMethod,
21
+ OutputFormat,
22
+ SliceMode,
23
+ )
24
+
25
+ Number = Union[int, float]
26
+
27
+
28
+ @dataclass
29
+ class ProcessingConfig:
30
+ """Every knob of the processing protocol, in the order it is applied.
31
+
32
+ The defaults are the recommended protocol for abdominal CT: reorient to
33
+ canonical RAS, clip to a soft-tissue window, resample to a common voxel
34
+ grid, mask to the region of interest, standardize the array shape.
35
+
36
+ Use :meth:`for_dataset` to get the values curated for a specific TCIA
37
+ collection (organs to segment, clipping window, output dimensions).
38
+ """
39
+
40
+ # ---- 1. orientation ------------------------------------------------
41
+ orient: bool = True
42
+ target_orientation: str = "RAS"
43
+
44
+ # ---- 2. segmentation ----------------------------------------------
45
+ segment: bool = False
46
+ organs: Optional[Sequence[str]] = None
47
+ totalsegmentator_task: str = "total"
48
+ remove_small_blobs: bool = True
49
+ fill_holes: bool = True
50
+ morphological_closing: bool = True
51
+ restrict_organs_to_tumor: bool = True
52
+ fast_segmentation: bool = False
53
+ device: Optional[str] = None
54
+ #: Where to keep the TotalSegmentator files. ``None`` writes them to a
55
+ #: temporary directory and deletes them once the masks are in memory; a
56
+ #: path keeps them, one subdirectory per series.
57
+ segmentation_dir: Optional[str] = None
58
+
59
+ # ---- 3. intensity clipping ----------------------------------------
60
+ clip: bool = True
61
+ clip_min: Optional[Number] = -200
62
+ clip_max: Optional[Number] = 300
63
+
64
+ # ---- 4. crop to content --------------------------------------------
65
+ #: Crop away the air around the body. Only useful once intensities are
66
+ #: clipped, since the threshold below separates tissue from air.
67
+ crop_to_content: bool = False
68
+ #: Voxels above this are kept. ``None`` uses `clip_min` when clipping is
69
+ #: on, and the image minimum otherwise.
70
+ crop_content_threshold: Optional[Number] = None
71
+ crop_content_padding: int = 5
72
+
73
+ # ---- 5. resampling -------------------------------------------------
74
+ resample: bool = True
75
+ target_spacing: Sequence[Optional[Number]] = (0.8, 0.8, 3.0)
76
+
77
+ # ---- 6. slice selection (2D only) ----------------------------------
78
+ dimensionality: Dimensionality = "3D"
79
+ #: ``"mask"`` keeps the slice with the most mask, ``"index"`` the slice
80
+ #: named by ``slice_index``.
81
+ slice_selection_mode: SliceMode = "mask"
82
+ #: Which mask value to measure. ``None`` uses the mask's only non-zero
83
+ #: value, which is unambiguous for a binary mask; a multi-label mask needs
84
+ #: this set (1 is organ and 2 tumor by this package's convention).
85
+ slice_selection_label: Optional[Union[int, Sequence[int]]] = None
86
+ slice_index: Optional[int] = None
87
+
88
+ # ---- 7. masking ----------------------------------------------------
89
+ mask: bool = True
90
+ mask_labels: Optional[Union[int, Sequence[int]]] = None
91
+ crop_to_mask: bool = True
92
+ crop_padding: int = 5
93
+
94
+ # ---- 8. size standardization ---------------------------------------
95
+ standardize_size: bool = True
96
+ target_shape: Optional[Sequence[Optional[int]]] = None
97
+ shape_percentile: float = 95.0
98
+
99
+ # ---- 9. intensity normalization -------------------------------------
100
+ normalize: bool = False
101
+ normalization_method: NormalizationMethod = "volume"
102
+ dataset_mean: Optional[float] = None
103
+ dataset_std: Optional[float] = None
104
+
105
+ # ---- output --------------------------------------------------------
106
+ output_format: OutputFormat = "nifti"
107
+ compress: bool = True
108
+ save_mask: bool = True
109
+ dtype: str = "float32"
110
+
111
+ # ---- radiomics -------------------------------------------------------
112
+ radiomics_labels: Sequence[int] = (1, 2)
113
+ radiomics_params: Optional[Union[str, dict]] = None
114
+
115
+ # ---- provenance ------------------------------------------------------
116
+ dataset: Optional[str] = None
117
+
118
+ def __post_init__(self) -> None:
119
+ self._normalize_sequences()
120
+ self.validate()
121
+
122
+ def _normalize_sequences(self) -> None:
123
+ """Make sequence fields tuples regardless of how they were supplied.
124
+
125
+ YAML has no tuple type, so without this a config would not compare
126
+ equal to itself after a save/load round trip.
127
+ """
128
+ self.target_spacing = tuple(self.target_spacing)
129
+ self.radiomics_labels = tuple(self.radiomics_labels)
130
+ if self.target_shape is not None:
131
+ self.target_shape = tuple(self.target_shape)
132
+ if self.organs is not None:
133
+ self.organs = list(self.organs)
134
+
135
+ # ------------------------------------------------------------------
136
+ # validation
137
+ # ------------------------------------------------------------------
138
+ def validate(self) -> "ProcessingConfig":
139
+ """Raise :class:`ValueError` on internally inconsistent settings."""
140
+ if self.dimensionality not in ("2D", "3D"):
141
+ raise ValueError(
142
+ f"dimensionality must be '2D' or '3D', got {self.dimensionality!r}"
143
+ )
144
+ if self.slice_selection_mode not in ("mask", "index"):
145
+ raise ValueError(
146
+ "slice_selection_mode must be 'mask' (the slice with the most mask) "
147
+ f"or 'index' (a slice number), got {self.slice_selection_mode!r}"
148
+ )
149
+ if (self.dimensionality == "2D" and self.slice_selection_mode == "index"
150
+ and self.slice_index is None):
151
+ raise ValueError(
152
+ "slice_selection_mode='index' requires slice_index=<slice number>. "
153
+ "Use slice_selection_mode='mask' to pick the slice with the most "
154
+ "mask instead."
155
+ )
156
+ if self.normalization_method not in ("volume", "dataset"):
157
+ raise ValueError(
158
+ "normalization_method must be 'volume' or 'dataset', got "
159
+ f"{self.normalization_method!r}"
160
+ )
161
+ if self.output_format not in OUTPUT_FORMATS:
162
+ raise ValueError(
163
+ f"output_format must be one of {', '.join(OUTPUT_FORMATS)}, "
164
+ f"got {self.output_format!r}. DICOM output is not supported: a "
165
+ "processed volume is no longer the acquisition its headers "
166
+ "describe."
167
+ )
168
+ if self.clip and self.clip_min is None and self.clip_max is None:
169
+ raise ValueError(
170
+ "clip=True requires clip_min and/or clip_max (e.g. -200/300 for a "
171
+ "soft-tissue window). Set clip=False to skip clipping."
172
+ )
173
+ if self.segment and not self.organs:
174
+ raise ValueError(
175
+ "segment=True requires `organs` (TotalSegmentator structure names, "
176
+ "e.g. ['kidney_left', 'kidney_right']). "
177
+ "See ProcessingConfig.for_dataset() for curated per-collection values."
178
+ )
179
+ if len(self.target_spacing) != 3:
180
+ raise ValueError(
181
+ f"target_spacing must have 3 entries, got {self.target_spacing!r}"
182
+ )
183
+ if self.target_shape is not None and len(self.target_shape) != 3:
184
+ raise ValueError(
185
+ f"target_shape must have 3 entries (or be None), got {self.target_shape!r}"
186
+ )
187
+ return self
188
+
189
+ # ------------------------------------------------------------------
190
+ # constructors
191
+ # ------------------------------------------------------------------
192
+ @classmethod
193
+ def resolve(cls, config: Any = None, **overrides: Any) -> "ProcessingConfig":
194
+ """A protocol from whatever names one.
195
+
196
+ ``None`` gives the defaults, a :class:`ProcessingConfig` is returned as
197
+ is, a mapping is expanded into one, and a string is either a path to a
198
+ saved protocol or the name of a collection whose curated protocol to
199
+ use::
200
+
201
+ ProcessingConfig.resolve("tcga-kirc")
202
+ ProcessingConfig.resolve("data/processed/processing_config.yaml")
203
+
204
+ This is what lets ``process()`` take any of them.
205
+ """
206
+ if config is None:
207
+ resolved = cls()
208
+ elif isinstance(config, ProcessingConfig):
209
+ resolved = config
210
+ elif isinstance(config, dict):
211
+ resolved = cls(**config)
212
+ elif isinstance(config, (str, os.PathLike)):
213
+ source = str(config)
214
+ resolved = cls.from_yaml(source) if os.path.exists(source) else cls.for_dataset(source)
215
+ else:
216
+ raise TypeError(
217
+ f"Cannot read a protocol from {type(config).__name__}. Expected a "
218
+ "ProcessingConfig, a collection name, a path to a saved protocol, "
219
+ "or None for the defaults."
220
+ )
221
+ return resolved.replace(**overrides) if overrides else resolved
222
+
223
+ @classmethod
224
+ def for_dataset(
225
+ cls, dataset: str, masked: bool = True, **overrides: Any
226
+ ) -> "ProcessingConfig":
227
+ """Build a config from the curated settings for a TCIA collection.
228
+
229
+ Parameters
230
+ ----------
231
+ dataset:
232
+ Key of :data:`~ctkit.constants.tcia_dataset_to_info`,
233
+ e.g. ``"tcga-kirc"``. Case- and separator-insensitive.
234
+ masked:
235
+ Whether the output will be masked to the ROI. Controls which set of
236
+ curated output dimensions is used.
237
+ **overrides:
238
+ Any field of :class:`ProcessingConfig`.
239
+ """
240
+ key = _normalize_dataset_key(dataset)
241
+ if key not in tcia_dataset_to_info:
242
+ raise KeyError(
243
+ f"Unknown dataset {dataset!r}. Known datasets: "
244
+ f"{', '.join(sorted(tcia_dataset_to_info))}"
245
+ )
246
+ info = tcia_dataset_to_info[key]
247
+
248
+ clip_min, clip_max = info.get("clip_min,clip_max", (None, None))
249
+ shape_key = "xdim,ydim,zdim_masked" if masked else "xdim,ydim,zdim_unmasked"
250
+ target_shape = info.get(shape_key)
251
+ if target_shape is not None and all(d is None for d in target_shape):
252
+ target_shape = None
253
+
254
+ organs = list(info.get("totalsegmentator_organs") or [])
255
+
256
+ params: dict = dict(
257
+ clip=clip_min is not None or clip_max is not None,
258
+ clip_min=clip_min,
259
+ clip_max=clip_max,
260
+ organs=organs or None,
261
+ segment=bool(organs),
262
+ totalsegmentator_task=info.get("totalsegmentator_task", "total"),
263
+ mask=masked,
264
+ target_shape=target_shape,
265
+ dataset=key,
266
+ )
267
+ params.update(overrides)
268
+ return cls(**params)
269
+
270
+ @classmethod
271
+ def radiomics(cls, dataset: Optional[str] = None, **overrides: Any) -> "ProcessingConfig":
272
+ """Preset for radiomic feature extraction.
273
+
274
+ PyRadiomics performs its own resampling, normalization and mask
275
+ handling, so those steps are disabled here to avoid applying them
276
+ twice. Orientation and clipping are kept.
277
+ """
278
+ params: dict = dict(
279
+ resample=False,
280
+ mask=False,
281
+ standardize_size=False,
282
+ normalize=False,
283
+ )
284
+ params.update(overrides)
285
+ if dataset is not None:
286
+ return cls.for_dataset(dataset, masked=False, **params)
287
+ return cls(**params)
288
+
289
+ @classmethod
290
+ def minimal(cls, **overrides: Any) -> "ProcessingConfig":
291
+ """Everything off: a no-op pipeline, useful as an ablation baseline."""
292
+ params: dict = dict(
293
+ orient=False,
294
+ segment=False,
295
+ clip=False,
296
+ resample=False,
297
+ mask=False,
298
+ standardize_size=False,
299
+ normalize=False,
300
+ )
301
+ params.update(overrides)
302
+ return cls(**params)
303
+
304
+ # ------------------------------------------------------------------
305
+ # serialization
306
+ # ------------------------------------------------------------------
307
+ def to_dict(self) -> dict:
308
+ return asdict(self)
309
+
310
+ @classmethod
311
+ def from_dict(cls, data: dict) -> "ProcessingConfig":
312
+ known = {f.name for f in fields(cls)}
313
+ unknown = set(data) - known
314
+ if unknown:
315
+ raise ValueError(
316
+ f"Unknown config keys: {', '.join(sorted(unknown))}. "
317
+ f"Valid keys: {', '.join(sorted(known))}"
318
+ )
319
+ return cls(**{k: v for k, v in data.items() if k in known})
320
+
321
+ def to_yaml(self, path: Optional[str] = None) -> str:
322
+ import yaml
323
+
324
+ text = yaml.safe_dump(self.to_dict(), sort_keys=False, default_flow_style=False)
325
+ if path:
326
+ with open(path, "w") as handle:
327
+ handle.write(text)
328
+ return text
329
+
330
+ @classmethod
331
+ def from_yaml(cls, path: str) -> "ProcessingConfig":
332
+ import yaml
333
+
334
+ with open(path) as handle:
335
+ return cls.from_dict(yaml.safe_load(handle) or {})
336
+
337
+ @classmethod
338
+ def load(cls, path: str) -> "ProcessingConfig":
339
+ """Load from a ``.yaml``/``.yml`` or ``.json`` file."""
340
+ if path.endswith(".json"):
341
+ with open(path) as handle:
342
+ return cls.from_dict(json.load(handle))
343
+ return cls.from_yaml(path)
344
+
345
+ def replace(self, **overrides: Any) -> "ProcessingConfig":
346
+ """Return a copy with `overrides` applied."""
347
+ data = self.to_dict()
348
+ data.update(overrides)
349
+ return type(self).from_dict(data)
350
+
351
+ # ------------------------------------------------------------------
352
+ # introspection
353
+ # ------------------------------------------------------------------
354
+ @property
355
+ def steps(self) -> list:
356
+ """Names of the steps that will actually run, in order."""
357
+ enabled = []
358
+ if self.orient:
359
+ enabled.append("orient")
360
+ if self.segment:
361
+ enabled.append("segment")
362
+ if self.clip:
363
+ enabled.append("clip")
364
+ if self.crop_to_content:
365
+ enabled.append("crop_to_content")
366
+ if self.resample:
367
+ enabled.append("resample")
368
+ if self.dimensionality == "2D":
369
+ enabled.append("select_slice")
370
+ if self.mask:
371
+ enabled.append("apply_mask")
372
+ if self.standardize_size:
373
+ enabled.append("standardize_size")
374
+ if self.normalize:
375
+ enabled.append("normalize")
376
+ return enabled
377
+
378
+ @property
379
+ def resolved_crop_threshold(self) -> Optional[Number]:
380
+ """The threshold :meth:`~ctkit.RadiologyImage.crop_to_content` will use.
381
+
382
+ ``None`` means the image minimum, which only separates body from air
383
+ once intensities have been clipped — so when clipping is on, the clip
384
+ minimum is used, exactly as it is for the masking and padding fills.
385
+ """
386
+ if self.crop_content_threshold is not None:
387
+ return self.crop_content_threshold
388
+ return self.clip_min if self.clip else None
389
+
390
+ @property
391
+ def needs_dataset_pass(self) -> bool:
392
+ """True when a step needs statistics pooled over the whole dataset."""
393
+ return bool(
394
+ (self.normalize and self.normalization_method == "dataset"
395
+ and self.dataset_mean is None)
396
+ or (self.standardize_size and self.target_shape is None)
397
+ )
398
+
399
+ def describe(self) -> str:
400
+ """Human-readable protocol summary, suitable for a methods section."""
401
+ lines = ["Processing protocol:"]
402
+ detail = {
403
+ "orient": f"reorient to canonical {self.target_orientation}",
404
+ "segment": (
405
+ f"segment {', '.join(self.organs or [])} with TotalSegmentator"
406
+ f" (task={self.totalsegmentator_task})"
407
+ ),
408
+ "clip": f"clip intensities to [{self.clip_min}, {self.clip_max}] HU",
409
+ "crop_to_content": (
410
+ "crop to the voxels above "
411
+ + ("the image minimum" if self.resolved_crop_threshold is None
412
+ else str(self.resolved_crop_threshold))
413
+ + f" with {self.crop_content_padding} voxel padding"
414
+ ),
415
+ "resample": f"resample to {tuple(self.target_spacing)} mm voxels",
416
+ "select_slice": (
417
+ f"keep axial slice {self.slice_index}"
418
+ if self.slice_selection_mode == "index"
419
+ else "keep the axial slice with the most mask"
420
+ + (
421
+ ""
422
+ if self.slice_selection_label is None
423
+ else f" (label {self.slice_selection_label})"
424
+ )
425
+ ),
426
+ "apply_mask": (
427
+ "mask to "
428
+ + ("all labels" if self.mask_labels is None else str(self.mask_labels))
429
+ + (f", crop to ROI with {self.crop_padding} voxel padding"
430
+ if self.crop_to_mask else "")
431
+ ),
432
+ "standardize_size": (
433
+ "crop/pad to "
434
+ + (str(tuple(self.target_shape)) if self.target_shape
435
+ else f"the {self.shape_percentile:g}th percentile shape of the dataset")
436
+ ),
437
+ "normalize": f"z-score normalize per {self.normalization_method}",
438
+ }
439
+ for index, step in enumerate(self.steps, start=1):
440
+ lines.append(f" {index}. {detail[step]}")
441
+ if len(lines) == 1:
442
+ lines.append(" (no steps enabled)")
443
+ return "\n".join(lines)
444
+
445
+
446
+ def _normalize_dataset_key(dataset: str) -> str:
447
+ """Map user spellings (``TCGA_KIRC``, ``TCGA-KIRC``) to registry keys."""
448
+ key = str(dataset).strip().lower().replace("_", "-").replace(" ", "-")
449
+ return key
450
+
451
+
452
+ # Backwards-compatible alias: the notebook parameter name.
453
+ DEFAULT_CONFIG = ProcessingConfig()