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/__init__.py +78 -0
- ct_seg/denoise/__init__.py +24 -0
- ct_seg/denoise/data.py +348 -0
- ct_seg/denoise/denoise_slice.py +190 -0
- ct_seg/denoise/denoise_volume.py +164 -0
- ct_seg/denoise/eval.py +94 -0
- ct_seg/denoise/loss.py +78 -0
- ct_seg/denoise/model.py +205 -0
- ct_seg/denoise/tiffs.py +105 -0
- ct_seg/denoise/train.py +385 -0
- ct_seg/denoise/utils.py +68 -0
- ct_seg/labeling.py +145 -0
- ct_seg/model.py +117 -0
- ct_seg/segment.py +461 -0
- ct_seg/som_bands.py +199 -0
- ct_seg/tracking.py +84 -0
- ct_seg/train.py +441 -0
- ct_seg/viewer.py +240 -0
- ct_segmentation_toolkit-0.2.0.dist-info/METADATA +156 -0
- ct_segmentation_toolkit-0.2.0.dist-info/RECORD +24 -0
- ct_segmentation_toolkit-0.2.0.dist-info/WHEEL +5 -0
- ct_segmentation_toolkit-0.2.0.dist-info/licenses/LICENSE +201 -0
- ct_segmentation_toolkit-0.2.0.dist-info/licenses/NOTICE +18 -0
- ct_segmentation_toolkit-0.2.0.dist-info/top_level.txt +1 -0
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()
|