ct-segmentation-toolkit 0.2.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.
ct_seg/train.py ADDED
@@ -0,0 +1,441 @@
1
+ """
2
+ Train U-Net for Supervised CT Segmentation
3
+ ===========================================
4
+ Trains a U-Net model on labeled CT slices for multi-class segmentation.
5
+
6
+ Usage:
7
+ python train_seg.py --images /path/to/tiffs --masks /path/to/mask_tiffs --num_classes 4
8
+ python train_seg.py --images /path/to/tiffs --masks /path/to/mask_tiffs --num_classes 4 --epochs 200 --batch_size 8
9
+
10
+ Mask format:
11
+ - Tiff images with integer pixel values 0, 1, 2, ... (num_classes - 1)
12
+ - Same filenames and dimensions as the corresponding image tiffs
13
+ - Can be created with label_tool.py (Napari-based)
14
+ """
15
+
16
+ import argparse
17
+ import json
18
+ import logging
19
+ import os
20
+ import sys
21
+ import time
22
+ from copy import deepcopy
23
+ from pathlib import Path
24
+
25
+ import numpy as np
26
+ import tifffile
27
+ import torch
28
+ import torch.nn as nn
29
+ from matplotlib import pyplot as plt
30
+ from torch.utils.data import DataLoader, Dataset
31
+ from tqdm import tqdm
32
+
33
+ from ct_seg import tracking
34
+ from ct_seg.model import UNetSegmentation, count_parameters
35
+
36
+ # =============================================================================
37
+ # Dataset
38
+ # =============================================================================
39
+
40
+
41
+ def natural_sorted(paths):
42
+ import re
43
+
44
+ def key(x):
45
+ return [int(c) if c.isdigit() else c for c in re.split(r"([0-9]+)", str(x))]
46
+
47
+ return sorted(paths, key=key)
48
+
49
+
50
+ class SegmentationDataset(Dataset):
51
+ """
52
+ Dataset for CT segmentation training.
53
+ Loads image-mask pairs from two directories of tiff files.
54
+ """
55
+
56
+ def __init__(self, image_dir, mask_dir, patch_size=256, augment=True, slice_range=None):
57
+ self.patch_size = patch_size
58
+ self.augment = augment
59
+
60
+ # Load images
61
+ img_paths = natural_sorted(list(Path(image_dir).glob("*.tif*")))
62
+ mask_paths = natural_sorted(list(Path(mask_dir).glob("*.tif*")))
63
+
64
+ if slice_range:
65
+ img_paths = img_paths[slice_range[0] : slice_range[1]]
66
+ mask_paths = mask_paths[slice_range[0] : slice_range[1]]
67
+
68
+ assert len(img_paths) == len(
69
+ mask_paths
70
+ ), f"Image count ({len(img_paths)}) != mask count ({len(mask_paths)})"
71
+
72
+ print(f"Loading {len(img_paths)} image-mask pairs...")
73
+ self.images = np.stack([tifffile.imread(str(p)) for p in tqdm(img_paths, desc="Images")])
74
+ self.masks = np.stack([tifffile.imread(str(p)) for p in tqdm(mask_paths, desc="Masks")])
75
+
76
+ # Normalize images
77
+ self.img_mean = self.images.mean()
78
+ self.img_std = self.images.std()
79
+ self.images = ((self.images - self.img_mean) / (self.img_std + 1e-8)).astype(np.float32)
80
+
81
+ # Ensure masks are integer class labels
82
+ self.masks = self.masks.astype(np.int64)
83
+
84
+ unique_classes = np.unique(self.masks)
85
+ print(f"Image shape: {self.images.shape}, Mask shape: {self.masks.shape}")
86
+ print(f"Classes found in masks: {unique_classes}")
87
+ print(f"Normalization: mean={self.img_mean:.6f}, std={self.img_std:.6f}")
88
+
89
+ def __len__(self):
90
+ return self.images.shape[0]
91
+
92
+ def __getitem__(self, idx):
93
+ img = self.images[idx]
94
+ mask = self.masks[idx]
95
+
96
+ # Random crop
97
+ h, w = img.shape
98
+ if h > self.patch_size and w > self.patch_size:
99
+ y = np.random.randint(0, h - self.patch_size)
100
+ x = np.random.randint(0, w - self.patch_size)
101
+ img = img[y : y + self.patch_size, x : x + self.patch_size]
102
+ mask = mask[y : y + self.patch_size, x : x + self.patch_size]
103
+
104
+ # Augmentation: random D4 symmetry (rotation + flip)
105
+ if self.augment:
106
+ k = np.random.randint(4)
107
+ img = np.rot90(img, k).copy()
108
+ mask = np.rot90(mask, k).copy()
109
+ if np.random.random() > 0.5:
110
+ img = np.fliplr(img).copy()
111
+ mask = np.fliplr(mask).copy()
112
+
113
+ # Add channel dimension [1, H, W]
114
+ img = img[np.newaxis]
115
+ return torch.from_numpy(img), torch.from_numpy(mask)
116
+
117
+
118
+ # =============================================================================
119
+ # Loss Functions
120
+ # =============================================================================
121
+
122
+
123
+ class DiceLoss(nn.Module):
124
+ """Soft Dice loss for multi-class segmentation."""
125
+
126
+ def __init__(self, num_classes, smooth=1.0):
127
+ super().__init__()
128
+ self.num_classes = num_classes
129
+ self.smooth = smooth
130
+
131
+ def forward(self, logits, targets):
132
+ # logits: [B, C, H, W], targets: [B, H, W] (integer class labels)
133
+ probs = torch.softmax(logits, dim=1) # [B, C, H, W]
134
+ targets_one_hot = torch.nn.functional.one_hot(targets, self.num_classes) # [B, H, W, C]
135
+ targets_one_hot = targets_one_hot.permute(0, 3, 1, 2).float() # [B, C, H, W]
136
+
137
+ dims = (0, 2, 3) # sum over batch, H, W
138
+ intersection = (probs * targets_one_hot).sum(dims)
139
+ union = probs.sum(dims) + targets_one_hot.sum(dims)
140
+
141
+ dice = (2.0 * intersection + self.smooth) / (union + self.smooth)
142
+ return 1.0 - dice.mean()
143
+
144
+
145
+ class CombinedLoss(nn.Module):
146
+ """Cross-Entropy + Dice loss."""
147
+
148
+ def __init__(self, num_classes, ce_weight=1.0, dice_weight=1.0):
149
+ super().__init__()
150
+ self.ce = nn.CrossEntropyLoss()
151
+ self.dice = DiceLoss(num_classes)
152
+ self.ce_weight = ce_weight
153
+ self.dice_weight = dice_weight
154
+
155
+ def forward(self, logits, targets):
156
+ return self.ce_weight * self.ce(logits, targets) + self.dice_weight * self.dice(
157
+ logits, targets
158
+ )
159
+
160
+
161
+ # =============================================================================
162
+ # Metrics
163
+ # =============================================================================
164
+
165
+
166
+ def compute_iou(pred, target, num_classes):
167
+ """Compute per-class IoU (Intersection over Union)."""
168
+ ious = []
169
+ for c in range(num_classes):
170
+ pred_c = pred == c
171
+ target_c = target == c
172
+ intersection = (pred_c & target_c).sum().item()
173
+ union = (pred_c | target_c).sum().item()
174
+ if union == 0:
175
+ ious.append(float("nan"))
176
+ else:
177
+ ious.append(intersection / union)
178
+ return ious
179
+
180
+
181
+ # =============================================================================
182
+ # Training
183
+ # =============================================================================
184
+
185
+
186
+ def main():
187
+ parser = argparse.ArgumentParser(description="Train U-Net Segmentation")
188
+ parser.add_argument("--images", required=True, help="Directory of image tiffs")
189
+ parser.add_argument(
190
+ "--masks", required=True, help="Directory of mask tiffs (integer class labels)"
191
+ )
192
+ parser.add_argument("--output", default="SegOutput", help="Output directory for checkpoints")
193
+ parser.add_argument(
194
+ "--num_classes", type=int, required=True, help="Number of segmentation classes"
195
+ )
196
+ parser.add_argument(
197
+ "--base_filters", type=int, default=32, help="Base filter count (default: 32)"
198
+ )
199
+ parser.add_argument(
200
+ "--patch_size", type=int, default=256, help="Training patch size (default: 256)"
201
+ )
202
+ parser.add_argument("--batch_size", type=int, default=8, help="Batch size (default: 8)")
203
+ parser.add_argument("--epochs", type=int, default=200, help="Max epochs (default: 200)")
204
+ parser.add_argument("--lr", type=float, default=0.001, help="Learning rate (default: 0.001)")
205
+ parser.add_argument(
206
+ "--val_split", type=float, default=0.15, help="Validation split fraction (default: 0.15)"
207
+ )
208
+ parser.add_argument("--slice_range", nargs=2, type=int, default=None, metavar=("START", "END"))
209
+ parser.add_argument("--device", default="cuda", help="Device (default: cuda)")
210
+ parser.add_argument(
211
+ "--no_mlflow", action="store_true", help="Disable MLflow experiment tracking"
212
+ )
213
+ args = parser.parse_args()
214
+
215
+ # Setup
216
+ os.makedirs(args.output, exist_ok=True)
217
+ os.makedirs(f"{args.output}/results", exist_ok=True)
218
+
219
+ logging.basicConfig(filename=f"{args.output}/training.log", level=logging.DEBUG)
220
+ logging.getLogger().addHandler(logging.StreamHandler(sys.stdout))
221
+
222
+ dev = torch.device(args.device if torch.cuda.is_available() else "cpu")
223
+ logging.info(f"Device: {dev}")
224
+
225
+ # Load dataset
226
+ full_dataset = SegmentationDataset(
227
+ args.images,
228
+ args.masks,
229
+ patch_size=args.patch_size,
230
+ augment=True,
231
+ slice_range=args.slice_range,
232
+ )
233
+
234
+ # Train/val split
235
+ n_val = max(1, int(len(full_dataset) * args.val_split))
236
+ n_train = len(full_dataset) - n_val
237
+ train_ds, val_ds = torch.utils.data.random_split(full_dataset, [n_train, n_val])
238
+ val_ds.dataset.augment = False # no augmentation for validation
239
+
240
+ train_dl = DataLoader(
241
+ train_ds,
242
+ batch_size=args.batch_size,
243
+ shuffle=True,
244
+ num_workers=4,
245
+ pin_memory=True,
246
+ prefetch_factor=2,
247
+ )
248
+ val_dl = DataLoader(
249
+ val_ds,
250
+ batch_size=args.batch_size,
251
+ shuffle=False,
252
+ num_workers=4,
253
+ pin_memory=True,
254
+ prefetch_factor=2,
255
+ )
256
+
257
+ logging.info(f"Train: {n_train} samples, Val: {n_val} samples")
258
+
259
+ # Model
260
+ model = UNetSegmentation(
261
+ in_channels=1, num_classes=args.num_classes, base_filters=args.base_filters
262
+ ).to(dev)
263
+ logging.info(f"Model parameters: {count_parameters(model):,}")
264
+
265
+ # Optimizer and loss
266
+ optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
267
+ scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
268
+ optimizer, mode="min", factor=0.5, patience=20
269
+ )
270
+ criterion = CombinedLoss(args.num_classes)
271
+
272
+ # Optional MLflow experiment tracking (no-op if mlflow absent or --no_mlflow).
273
+ track = not args.no_mlflow
274
+ tracking.start(run_name=f"unet_c{args.num_classes}", enabled=track)
275
+ tracking.log_params(
276
+ {
277
+ "num_classes": args.num_classes,
278
+ "base_filters": args.base_filters,
279
+ "patch_size": args.patch_size,
280
+ "batch_size": args.batch_size,
281
+ "epochs": args.epochs,
282
+ "lr": args.lr,
283
+ "val_split": args.val_split,
284
+ "device": str(dev),
285
+ "train_samples": n_train,
286
+ "val_samples": n_val,
287
+ },
288
+ enabled=track,
289
+ )
290
+
291
+ # Training loop
292
+ best_val_loss = float("inf")
293
+ best_iou = 0
294
+ train_losses, val_losses = [], []
295
+ start_time = time.time()
296
+
297
+ for epoch in range(1, args.epochs + 1):
298
+ tick = time.time()
299
+
300
+ # Train
301
+ model.train()
302
+ epoch_loss = []
303
+ for imgs, masks in train_dl:
304
+ imgs, masks = imgs.to(dev), masks.to(dev)
305
+ optimizer.zero_grad()
306
+ logits = model(imgs)
307
+ loss = criterion(logits, masks)
308
+ loss.backward()
309
+ optimizer.step()
310
+ epoch_loss.append(loss.item())
311
+
312
+ train_loss = np.mean(epoch_loss)
313
+
314
+ # Validate
315
+ model.eval()
316
+ val_epoch_loss = []
317
+ all_ious = []
318
+ with torch.no_grad():
319
+ for imgs, masks in val_dl:
320
+ imgs, masks = imgs.to(dev), masks.to(dev)
321
+ logits = model(imgs)
322
+ loss = criterion(logits, masks)
323
+ val_epoch_loss.append(loss.item())
324
+
325
+ pred = torch.argmax(logits, dim=1)
326
+ ious = compute_iou(pred.cpu().numpy(), masks.cpu().numpy(), args.num_classes)
327
+ all_ious.append(ious)
328
+
329
+ val_loss = np.mean(val_epoch_loss)
330
+ mean_iou = np.nanmean(all_ious)
331
+
332
+ scheduler.step(val_loss)
333
+
334
+ train_losses.append(train_loss)
335
+ val_losses.append(val_loss)
336
+
337
+ tracking.log_metrics(
338
+ {
339
+ "train_loss": float(train_loss),
340
+ "val_loss": float(val_loss),
341
+ "mean_iou": float(mean_iou),
342
+ "lr": optimizer.param_groups[0]["lr"],
343
+ },
344
+ step=epoch,
345
+ enabled=track,
346
+ )
347
+
348
+ ep_time = time.time() - tick
349
+ logging.info(
350
+ f"Epoch {epoch}/{args.epochs} | "
351
+ f"Train: {train_loss:.4f} | Val: {val_loss:.4f} | "
352
+ f"mIoU: {mean_iou:.4f} | LR: {optimizer.param_groups[0]['lr']:.6f} | "
353
+ f"{ep_time:.1f}s"
354
+ )
355
+
356
+ # Save best model (by val loss)
357
+ if val_loss < best_val_loss:
358
+ best_val_loss = val_loss
359
+ torch.save(
360
+ {
361
+ "model_state_dict": deepcopy(model.state_dict()),
362
+ "optimizer_state_dict": deepcopy(optimizer.state_dict()),
363
+ "epoch": epoch,
364
+ "val_loss": val_loss,
365
+ "num_classes": args.num_classes,
366
+ "base_filters": args.base_filters,
367
+ "img_mean": full_dataset.img_mean,
368
+ "img_std": full_dataset.img_std,
369
+ },
370
+ f"{args.output}/best_model.pth",
371
+ )
372
+ logging.info(f" -> Saved best model (val_loss={val_loss:.4f})")
373
+
374
+ # Save best model (by mIoU)
375
+ if mean_iou > best_iou:
376
+ best_iou = mean_iou
377
+ torch.save(
378
+ {
379
+ "model_state_dict": deepcopy(model.state_dict()),
380
+ "optimizer_state_dict": deepcopy(optimizer.state_dict()),
381
+ "epoch": epoch,
382
+ "mean_iou": mean_iou,
383
+ "num_classes": args.num_classes,
384
+ "base_filters": args.base_filters,
385
+ "img_mean": full_dataset.img_mean,
386
+ "img_std": full_dataset.img_std,
387
+ },
388
+ f"{args.output}/best_iou_model.pth",
389
+ )
390
+ logging.info(f" -> Saved best IoU model (mIoU={mean_iou:.4f})")
391
+
392
+ # Plot training curves every 10 epochs
393
+ if epoch % 10 == 0:
394
+ plt.figure(figsize=(10, 5))
395
+ plt.plot(train_losses, label="Train Loss")
396
+ plt.plot(val_losses, label="Val Loss")
397
+ plt.xlabel("Epoch")
398
+ plt.ylabel("Loss (CE + Dice)")
399
+ plt.title("Segmentation Training Progress")
400
+ plt.legend()
401
+ plt.savefig(f"{args.output}/results/training_curve.png", dpi=150)
402
+ plt.close()
403
+
404
+ total_time = time.time() - start_time
405
+ logging.info(f"\nTraining complete in {total_time:.0f}s")
406
+ logging.info(f"Best val loss: {best_val_loss:.4f}")
407
+ logging.info(f"Best mIoU: {best_iou:.4f}")
408
+
409
+ # Save training metadata
410
+ meta = {
411
+ "num_classes": args.num_classes,
412
+ "base_filters": args.base_filters,
413
+ "patch_size": args.patch_size,
414
+ "batch_size": args.batch_size,
415
+ "epochs": args.epochs,
416
+ "lr": args.lr,
417
+ "best_val_loss": float(best_val_loss),
418
+ "best_iou": float(best_iou),
419
+ "training_time_s": total_time,
420
+ "img_mean": float(full_dataset.img_mean),
421
+ "img_std": float(full_dataset.img_std),
422
+ }
423
+ with open(f"{args.output}/training_info.json", "w") as f:
424
+ json.dump(meta, f, indent=2)
425
+
426
+ # Log final summary + best checkpoint to MLflow, then close the run.
427
+ tracking.log_metrics(
428
+ {
429
+ "best_val_loss": float(best_val_loss),
430
+ "best_iou": float(best_iou),
431
+ "training_time_s": float(total_time),
432
+ },
433
+ enabled=track,
434
+ )
435
+ tracking.log_artifact(f"{args.output}/best_model.pth", enabled=track)
436
+ tracking.log_artifact(f"{args.output}/training_info.json", enabled=track)
437
+ tracking.end(enabled=track)
438
+
439
+
440
+ if __name__ == "__main__":
441
+ main()
ct_seg/viewer.py ADDED
@@ -0,0 +1,240 @@
1
+ """
2
+ 3D Volume Viewer for CT Segmentation Results
3
+ =============================================
4
+ Two viewing modes:
5
+ 1. Orthogonal slice viewer (Napari) - interactive XY/XZ/YZ navigation
6
+ 2. 3D volume rendering (PyVista) - rotatable 3D view of segmented phases
7
+
8
+ Usage:
9
+ # Orthogonal viewer (Napari)
10
+ python viewer.py --images /path/to/tiffs --mode ortho
11
+ python viewer.py --images /path/to/tiffs --labels /path/to/masks --mode ortho
12
+
13
+ # 3D volume rendering (PyVista)
14
+ python viewer.py --labels /path/to/masks --mode 3d --num_classes 4
15
+ python viewer.py --images /path/to/tiffs --labels /path/to/masks --mode 3d --num_classes 4
16
+
17
+ # Both
18
+ python viewer.py --images /path/to/tiffs --labels /path/to/masks --mode both
19
+
20
+ Options:
21
+ --slice_range 500 600 Load only a subset of slices
22
+ --downsample 2 Downsample volume for faster 3D rendering
23
+ """
24
+
25
+ import argparse
26
+ import re
27
+ import sys
28
+ from pathlib import Path
29
+
30
+ import numpy as np
31
+ import tifffile
32
+
33
+
34
+ def natural_sorted(paths):
35
+ def key(x):
36
+ return [int(c) if c.isdigit() else c for c in re.split(r"([0-9]+)", str(x))]
37
+
38
+ return sorted(paths, key=key)
39
+
40
+
41
+ def load_stack(dir_path, start=None, end=None, downsample=1):
42
+ from PIL import Image
43
+
44
+ dir_path = Path(dir_path)
45
+ paths = []
46
+ for ext in ["*.tif*", "*.png", "*.jpg", "*.jpeg"]:
47
+ paths.extend(dir_path.glob(ext))
48
+ paths = natural_sorted(list(set(paths)))
49
+ if start is not None and end is not None:
50
+ paths = paths[start:end]
51
+ if not paths:
52
+ raise FileNotFoundError(f"No image files in {dir_path}")
53
+ from tqdm import tqdm
54
+
55
+ # Use tifffile for tiffs (preserves raw values), PIL for others
56
+ def read_img(p):
57
+ if p.suffix.lower() in [".tif", ".tiff"]:
58
+ return tifffile.imread(str(p))
59
+ else:
60
+ return np.array(Image.open(str(p)).convert("L"))
61
+
62
+ imgs = np.stack([read_img(p) for p in tqdm(paths, desc=f"Loading {dir_path.name}")])
63
+ if downsample > 1:
64
+ imgs = imgs[::downsample, ::downsample, ::downsample]
65
+ return imgs
66
+
67
+
68
+ # =============================================================================
69
+ # Orthogonal Viewer (Napari)
70
+ # =============================================================================
71
+
72
+
73
+ def view_ortho(images=None, labels=None):
74
+ """Open Napari with orthogonal slice viewer."""
75
+ try:
76
+ import napari
77
+ except ImportError:
78
+ print("Napari required: pip install 'napari[all]'")
79
+ sys.exit(1)
80
+
81
+ viewer = napari.Viewer(title="CT Volume - Orthogonal Viewer", ndisplay=3)
82
+
83
+ if images is not None:
84
+ p2, p98 = np.percentile(images, (2, 98))
85
+ viewer.add_image(
86
+ np.clip(images, p2, p98), name="CT Volume", colormap="gray", rendering="mip"
87
+ )
88
+
89
+ if labels is not None:
90
+ label_colors = {
91
+ 0: [0, 0, 0, 0],
92
+ 1: [1, 0.2, 0.2, 0.6],
93
+ 2: [0.2, 1, 0.2, 0.6],
94
+ 3: [0.2, 0.4, 1, 0.6],
95
+ 4: [1, 1, 0.2, 0.6],
96
+ 5: [1, 0.2, 1, 0.6],
97
+ 6: [0.2, 1, 1, 0.6],
98
+ }
99
+ viewer.add_labels(labels, name="Segmentation", color=label_colors)
100
+
101
+ print("\nNapari Controls:")
102
+ print(" - Toggle 2D/3D view: button in bottom-left or press Ctrl+Y")
103
+ print(" - Scroll: navigate slices")
104
+ print(" - Click layer to toggle visibility")
105
+ print(" - Right-click layer for options")
106
+
107
+ napari.run()
108
+
109
+
110
+ # =============================================================================
111
+ # 3D Volume Rendering (PyVista)
112
+ # =============================================================================
113
+
114
+
115
+ def view_3d(labels, num_classes, images=None, downsample=1):
116
+ """Render 3D volume of segmented phases using PyVista."""
117
+ try:
118
+ import pyvista as pv
119
+ except ImportError:
120
+ print("PyVista required: pip install pyvista")
121
+ sys.exit(1)
122
+
123
+ # Colors for each class
124
+ class_colors = [
125
+ [0.1, 0.1, 0.1], # class 0: dark (background, often hidden)
126
+ [0.9, 0.2, 0.2], # class 1: red
127
+ [0.2, 0.9, 0.2], # class 2: green
128
+ [0.2, 0.4, 0.9], # class 3: blue
129
+ [0.9, 0.9, 0.2], # class 4: yellow
130
+ [0.9, 0.2, 0.9], # class 5: magenta
131
+ [0.2, 0.9, 0.9], # class 6: cyan
132
+ ]
133
+
134
+ class_names = [f"Class {i}" for i in range(num_classes)]
135
+
136
+ # Create plotter
137
+ plotter = pv.Plotter(title="3D Segmentation Viewer")
138
+
139
+ # Add each class as a separate surface
140
+ for c in range(1, num_classes): # skip class 0 (usually background)
141
+ binary = (labels == c).astype(np.uint8)
142
+
143
+ if binary.sum() == 0:
144
+ print(f"Class {c}: no voxels, skipping")
145
+ continue
146
+
147
+ # Create a uniform grid
148
+ grid = pv.ImageData(dimensions=np.array(binary.shape) + 1)
149
+ grid.cell_data["values"] = binary.flatten(order="F")
150
+
151
+ # Extract surface via threshold
152
+ surface = grid.threshold(0.5)
153
+
154
+ if surface.n_cells > 0:
155
+ color = class_colors[c] if c < len(class_colors) else [0.5, 0.5, 0.5]
156
+ voxel_count = binary.sum()
157
+ label = f"{class_names[c]} ({voxel_count:,} voxels)"
158
+ plotter.add_mesh(surface, color=color, opacity=0.6, label=label, smooth_shading=True)
159
+ print(f"Added {label}")
160
+
161
+ plotter.add_legend()
162
+ plotter.add_axes()
163
+ plotter.show_grid()
164
+ plotter.camera_position = "iso"
165
+
166
+ print("\nPyVista Controls:")
167
+ print(" - Left-click + drag: rotate")
168
+ print(" - Right-click + drag: zoom")
169
+ print(" - Middle-click + drag: pan")
170
+ print(" - Q: quit")
171
+
172
+ plotter.show()
173
+
174
+
175
+ # =============================================================================
176
+ # Main
177
+ # =============================================================================
178
+
179
+
180
+ def main():
181
+ parser = argparse.ArgumentParser(description="3D CT Volume Viewer")
182
+ parser.add_argument("--images", default=None, help="Directory of image tiffs")
183
+ parser.add_argument("--labels", default=None, help="Directory of segmentation mask tiffs")
184
+ parser.add_argument(
185
+ "--mode",
186
+ required=True,
187
+ choices=["ortho", "3d", "both"],
188
+ help="Viewing mode: ortho (Napari), 3d (PyVista), or both",
189
+ )
190
+ parser.add_argument(
191
+ "--num_classes", type=int, default=4, help="Number of classes (for 3D view)"
192
+ )
193
+ parser.add_argument("--slice_range", nargs=2, type=int, default=None, metavar=("START", "END"))
194
+ parser.add_argument(
195
+ "--downsample", type=int, default=1, help="Downsample factor for 3D rendering (default: 1)"
196
+ )
197
+ args = parser.parse_args()
198
+
199
+ if args.images is None and args.labels is None:
200
+ parser.error("At least one of --images or --labels is required")
201
+
202
+ start = args.slice_range[0] if args.slice_range else None
203
+ end = args.slice_range[1] if args.slice_range else None
204
+
205
+ images = None
206
+ labels = None
207
+
208
+ if args.images:
209
+ ds = args.downsample if args.mode in ["3d", "both"] else 1
210
+ images = load_stack(args.images, start=start, end=end, downsample=ds)
211
+ print(f"Images: {images.shape} ({images.dtype})")
212
+
213
+ if args.labels:
214
+ ds = args.downsample if args.mode in ["3d", "both"] else 1
215
+ labels = load_stack(args.labels, start=start, end=end, downsample=ds)
216
+ print(f"Labels: {labels.shape} ({labels.dtype})")
217
+ print(f"Classes: {np.unique(labels)}")
218
+
219
+ if args.mode == "ortho":
220
+ view_ortho(images, labels)
221
+ elif args.mode == "3d":
222
+ if labels is None:
223
+ parser.error("--labels is required for 3D mode")
224
+ view_3d(labels, args.num_classes, images, args.downsample)
225
+ elif args.mode == "both":
226
+ if labels is None:
227
+ parser.error("--labels is required for 3D mode")
228
+ print("\nOpening 3D viewer first (close to continue to ortho viewer)...")
229
+ view_3d(labels, args.num_classes, images, args.downsample)
230
+ print("\nOpening orthogonal viewer...")
231
+ # Reload without downsample for ortho
232
+ if args.images:
233
+ images = load_stack(args.images, start=start, end=end)
234
+ if args.labels:
235
+ labels = load_stack(args.labels, start=start, end=end)
236
+ view_ortho(images, labels)
237
+
238
+
239
+ if __name__ == "__main__":
240
+ main()