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.
@@ -0,0 +1,385 @@
1
+ """
2
+ main.py — Training entry point for 2.5D Noise2Inverse (N2I).
3
+
4
+ Trains a self-supervised U-Net that denoises computed-tomography (CT)
5
+ reconstructions without any ground-truth (clean) images, following the
6
+ Noise2Inverse framework. Two statistically independent sub-reconstructions
7
+ (e.g. reconstructions built from even- vs. odd-numbered projections) supervise
8
+ each other: the network learns to predict one split from the other. Because the
9
+ noise is independent between the splits while the underlying object is shared,
10
+ the optimum of this objective is the denoised signal. The "2.5D" variant feeds
11
+ five adjacent axial slices as input channels so the model can use through-plane
12
+ context while still predicting a single 2D slice.
13
+
14
+ Inputs : a YAML config pointing at the reconstruction directory (see
15
+ baseline_config.yaml) plus the even/odd sub-reconstruction TIFF stacks.
16
+ Outputs : a TrainOutput/ directory (created next to the reconstructions) holding
17
+ the best checkpoints, a training log, and periodic preview PNGs.
18
+
19
+ Training procedure:
20
+ * Warm up using a pixel-wise L1 loss only.
21
+ * After the LCL warm-up, add a Laplacian Contrast Loss term (loss.py) that
22
+ steers contrast toward edges to keep them sharp.
23
+ * Each epoch, keep three "best" checkpoints: lowest validation L1 loss,
24
+ lowest LCL loss, and highest Laplacian edge score (eval.py).
25
+
26
+ Usage:
27
+ python main.py -gpus=0 -config=/path/to/config.yaml
28
+ (normally launched via train.sh)
29
+
30
+ Author: Cameron Renteria <crentb23@gmail.com>
31
+ License: Apache-2.0 (see LICENSE)
32
+ """
33
+
34
+ import argparse
35
+ import logging
36
+ import os
37
+ import shutil
38
+ import sys
39
+ import time
40
+ from copy import deepcopy
41
+
42
+ import numpy as np
43
+ import torch
44
+ import yaml
45
+ from matplotlib import pyplot as plt
46
+ from torch.utils.data import DataLoader
47
+
48
+ from ct_seg.denoise.data import TomoDatasetTrain
49
+ from ct_seg.denoise.eval import laplacian_score_batch
50
+ from ct_seg.denoise.loss import LCL
51
+ from ct_seg.denoise.model import unet_ns_gn
52
+ from ct_seg.denoise.utils import save2img
53
+
54
+
55
+ def count_parameters(model):
56
+ """Return the number of trainable (gradient-requiring) parameters in `model`."""
57
+ return sum(p.numel() for p in model.parameters() if p.requires_grad)
58
+
59
+
60
+ def main(args):
61
+
62
+ # read the YAML file
63
+ with open(args.config, "r") as file:
64
+ params = yaml.safe_load(file)
65
+
66
+ START_TIME = time.time()
67
+
68
+ # create directory containing training results in the directory of the reconstructions
69
+ path_to_reconstructions = params["dataset"]["directory_to_reconstructions"]
70
+ odir = path_to_reconstructions + "/" + "TrainOutput"
71
+ if os.path.isdir(odir):
72
+ shutil.rmtree(odir)
73
+ os.mkdir(odir)
74
+ os.mkdir(f"{odir}/results")
75
+
76
+ # create output log
77
+ logging.basicConfig(filename=f"{odir}/Noise2Inverse.log", level=logging.DEBUG)
78
+ logging.getLogger().addHandler(logging.StreamHandler(sys.stdout))
79
+
80
+ # setup device
81
+ dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
82
+ if dev.type == "cuda":
83
+ torch.cuda.set_device(0)
84
+ logging.info(f"Using device: {dev}")
85
+
86
+ logging.info("\nLoading data into CPU memory, it will take a while ... ...")
87
+ ds_train = TomoDatasetTrain(params=params, config_file=args.config)
88
+ dl_train = DataLoader(
89
+ dataset=ds_train,
90
+ batch_size=params["train"]["mbsz"],
91
+ shuffle=True,
92
+ num_workers=4,
93
+ drop_last=False,
94
+ prefetch_factor=2,
95
+ pin_memory=True,
96
+ )
97
+
98
+ logging.info(
99
+ f"\nLoaded %d samples, {ds_train.samples}, into CPU memory for training." % (len(ds_train),)
100
+ )
101
+
102
+ # initialized model from scratch
103
+ model = unet_ns_gn(ich=5, start_filter_size=16, channels_per_group=8).to(dev)
104
+ optimizer = torch.optim.Adam(model.parameters(), lr=params["train"]["lr"])
105
+ logging.info("\nInitializing model from scratch\n")
106
+
107
+ print(f"Number of model parameters: {count_parameters(model):,}")
108
+ model_updates = 0
109
+
110
+ # Loss functions and warm-up schedule:
111
+ # criterion : pixel-wise L1 loss (used for the entire run)
112
+ # criterion_lcl : Laplacian Contrast Loss, the edge-sharpening term (loss.py)
113
+ # beta : weight applied to the LCL term once it is switched on
114
+ # lcl_warmup : number of model updates of pure-L1 warm-up before LCL kicks in
115
+ criterion = torch.nn.L1Loss()
116
+ criterion_lcl = LCL()
117
+ beta = 0.01
118
+ lcl_warmup = 1500
119
+ warmup = params["train"]["warmup"]
120
+ continue_warmup = True
121
+
122
+ train_loss, val_loss = [], []
123
+ edge_values = []
124
+ train_lcl_loss, val_lcl_loss = [], []
125
+
126
+ best_val_loss, best_edge, best_lcl_loss = np.inf, 0, np.inf
127
+ best_val_epoch, best_edge_epoch, best_lcl_epoch = 0, 0, 0
128
+
129
+ # start training
130
+ for epoch in range(1, params["train"]["maxep"] + 1):
131
+
132
+ step_losses, step_val_losses, step_lcl_loss, step_lcl_val_loss, step_edge_values = (
133
+ [],
134
+ [],
135
+ [],
136
+ [],
137
+ [],
138
+ )
139
+
140
+ tick_ep = time.time()
141
+
142
+ model.train()
143
+ # training loop
144
+ for X_mb, Y_mb in dl_train:
145
+ optimizer.zero_grad()
146
+
147
+ X_mb_dev = X_mb.to(dev)
148
+ Y_mb_dev = Y_mb.to(dev)
149
+
150
+ # Symmetric Noise2Inverse update. X_mb and Y_mb are the two independent
151
+ # sub-reconstruction splits. Each step trains the network in BOTH
152
+ # directions -- predict the center slice (channel index 2) of one split
153
+ # from the other -- so neither split is privileged as 'input' vs 'target'.
154
+ # While model_updates <= lcl_warmup only the L1 term is optimized;
155
+ # afterwards the edge-aware LCL term (loss.py) is added in.
156
+ if model_updates <= lcl_warmup:
157
+ optimizer.zero_grad()
158
+
159
+ # Process first view
160
+ pred_view1 = model(X_mb_dev)
161
+
162
+ loss_view1 = criterion(pred_view1.squeeze(dim=1), Y_mb_dev[:, 2])
163
+
164
+ loss_view1.backward()
165
+ optimizer.step()
166
+
167
+ optimizer.zero_grad()
168
+
169
+ # Process second view
170
+ pred_view2 = model(Y_mb_dev)
171
+ loss_view2 = criterion(pred_view2.squeeze(dim=1), X_mb_dev[:, 2])
172
+
173
+ loss_view2.backward()
174
+ optimizer.step()
175
+
176
+ loss_lcl1 = torch.tensor(0.0)
177
+ loss_lcl2 = torch.tensor(0.0)
178
+
179
+ else:
180
+ optimizer.zero_grad()
181
+
182
+ # Process first view
183
+ pred_view1 = model(X_mb_dev)
184
+
185
+ loss_view1 = criterion(pred_view1.squeeze(dim=1), Y_mb_dev[:, 2])
186
+ loss_lcl1 = criterion_lcl(pred_view1) * beta
187
+ loss_total1 = loss_view1 + loss_lcl1
188
+
189
+ loss_total1.backward()
190
+ optimizer.step()
191
+
192
+ optimizer.zero_grad()
193
+
194
+ # Process second view
195
+ pred_view2 = model(Y_mb_dev)
196
+ loss_view2 = criterion(pred_view2.squeeze(dim=1), X_mb_dev[:, 2])
197
+ loss_lcl2 = criterion_lcl(pred_view2) * beta
198
+ loss_total2 = loss_view2 + loss_lcl2
199
+
200
+ loss_total2.backward()
201
+ optimizer.step()
202
+
203
+ loss = loss_view1 + loss_view2
204
+ step_losses.append(loss.detach().cpu().numpy())
205
+ loss_lcl = loss_lcl1 + loss_lcl2
206
+ step_lcl_loss.append(loss_lcl.detach().cpu().numpy())
207
+ model_updates += 1
208
+
209
+ model.eval()
210
+ with torch.no_grad():
211
+ # validation loop
212
+ for X_mb, Y_mb in dl_train:
213
+ X_mb_dev = X_mb.to(dev)
214
+ Y_mb_dev = Y_mb.to(dev)
215
+
216
+ pred_view1 = model(X_mb_dev)
217
+ loss_view1 = criterion(pred_view1.squeeze(dim=1), Y_mb_dev[:, 2])
218
+ loss_lcl1 = criterion_lcl(pred_view1) * beta
219
+
220
+ pred_view2 = model(Y_mb_dev)
221
+ loss_view2 = criterion(pred_view2.squeeze(dim=1), X_mb_dev[:, 2])
222
+ loss_lcl2 = criterion_lcl(pred_view2) * beta
223
+
224
+ loss = loss_view1 + loss_view2
225
+
226
+ lap_score = laplacian_score_batch(pred_view1.cpu()) + laplacian_score_batch(
227
+ pred_view2.cpu()
228
+ )
229
+ step_edge_values.append(lap_score)
230
+
231
+ step_val_losses.append(loss.detach().cpu().numpy())
232
+ step_lcl_val_loss.append(loss_lcl.cpu().numpy())
233
+
234
+ ep_time = time.time() - tick_ep
235
+ logging.info(f"\nEpoch {epoch}")
236
+ iter_prints = f"[Train] L1 loss: {np.mean(step_losses):.6f}, {step_losses[0]:.6f} => {step_losses[-1]:.6f}, rate: {ep_time:.2f}s/ep"
237
+ logging.info(iter_prints)
238
+ iter_prints = f"[Train] LCL loss: {np.mean(step_lcl_loss):.6f}, {step_lcl_loss[0]:.6f} => {step_lcl_loss[-1]:.6f}, rate: {ep_time:.2f}s/ep"
239
+ logging.info(iter_prints)
240
+ iter_prints = f"[Val] L1 loss: {np.mean(step_val_losses):.6f}, {step_val_losses[0]:.6f} => {step_val_losses[-1]:.6f}, rate: {ep_time:.2f}s/ep"
241
+ logging.info(iter_prints)
242
+ iter_prints = f"[Val] LCL loss: {np.mean(step_lcl_val_loss):.6f}, {step_lcl_val_loss[0]:.6f} => {step_lcl_val_loss[-1]:.6f}, rate: {ep_time:.2f}s/ep"
243
+ logging.info(iter_prints)
244
+ iter_prints = f"[Val] EDGE Value: {np.mean(step_edge_values):.4f}, {step_edge_values[0]:.4f} => {step_edge_values[-1]:.4f}, rate: {ep_time:.2f}s/ep"
245
+ logging.info(iter_prints)
246
+
247
+ train_loss.append(np.mean(step_losses))
248
+ val_loss.append(np.mean(step_val_losses))
249
+ train_lcl_loss.append(np.mean(step_lcl_loss))
250
+ val_lcl_loss.append(np.mean(step_lcl_val_loss))
251
+ edge_values.append(np.mean(step_edge_values))
252
+
253
+ # Save the best model with the lowest lcl loss
254
+ if np.mean(step_lcl_val_loss) < best_lcl_loss:
255
+ best_lcl_loss = np.mean(step_lcl_val_loss)
256
+ best_lcl_epoch = epoch
257
+ mdl_fname = f"{odir}/best_lcl_model.pth"
258
+ torch.save(
259
+ {
260
+ "model_state_dict": deepcopy(model.state_dict()),
261
+ "optimizer_state_dict": deepcopy(optimizer.state_dict()),
262
+ },
263
+ mdl_fname,
264
+ )
265
+
266
+ # Save the best model with the lowest val loss
267
+ if np.mean(step_val_losses) < best_val_loss:
268
+ best_val_loss = np.mean(step_val_losses)
269
+ best_val_epoch = epoch
270
+ mdl_fname = f"{odir}/best_val_model.pth"
271
+ torch.save(
272
+ {
273
+ "model_state_dict": deepcopy(model.state_dict()),
274
+ "optimizer_state_dict": deepcopy(optimizer.state_dict()),
275
+ },
276
+ mdl_fname,
277
+ )
278
+
279
+ # Save the best model with the highest edge value
280
+ if np.mean(step_edge_values) > best_edge:
281
+ best_edge = np.mean(step_edge_values)
282
+ best_edge_epoch = epoch
283
+ mdl_fname = f"{odir}/best_edge_model.pth"
284
+ torch.save(
285
+ {
286
+ "model_state_dict": deepcopy(model.state_dict()),
287
+ "optimizer_state_dict": deepcopy(optimizer.state_dict()),
288
+ },
289
+ mdl_fname,
290
+ )
291
+
292
+ # End of warm-up: once BOTH warm-up thresholds (warmup and lcl_warmup) are
293
+ # passed, reset the best-metric trackers so that checkpoints selected during
294
+ # the pre-LCL phase do not dominate the rest of training.
295
+ if model_updates > warmup and model_updates > lcl_warmup and continue_warmup:
296
+
297
+ best_edge, best_lcl_loss = 0, np.inf
298
+ best_edge_epoch, best_lcl_epoch = 0, 0
299
+
300
+ continue_warmup = False
301
+
302
+ CRNT_TIME = time.time()
303
+ logging.info(f"[Info] Training Time: {CRNT_TIME-START_TIME:.2f} seconds")
304
+
305
+ # option to view the denoising process during training
306
+ if epoch % 5 == 0:
307
+ ridx = np.random.randint(pred_view1.shape[0])
308
+
309
+ save2img(
310
+ pred_view1[ridx, -1].detach().cpu().numpy(),
311
+ "%s/results/_%d_pred_view1.png" % (odir, epoch),
312
+ )
313
+ save2img(
314
+ pred_view2[ridx, -1].detach().cpu().numpy(),
315
+ "%s/results/_%d_pred_view2.png" % (odir, epoch),
316
+ )
317
+ save2img(
318
+ X_mb_dev[ridx, 2].detach().cpu().numpy(), "%s/results/_%d_gt.png" % (odir, epoch)
319
+ )
320
+ save2img(
321
+ Y_mb_dev[ridx, 2].detach().cpu().numpy(), "%s/results/_%d_noise.png" % (odir, epoch)
322
+ )
323
+
324
+ # Keep track of when/where the best model is
325
+ logging.info(f"Lowest model validation loss {best_val_loss:.6f} at epoch {best_val_epoch}")
326
+ logging.info(f"Lowest model LCL loss {best_lcl_loss:.6f} at epoch {best_lcl_epoch}")
327
+ logging.info(f"Highest model EDGE score {best_edge:.6f} at epoch {best_edge_epoch}")
328
+ logging.info(f"Number of model updates: {model_updates:,}")
329
+ logging.info(f"Is model warming up?: {continue_warmup}")
330
+
331
+ # View the training/validation loss during training
332
+ if epoch % 5 == 0:
333
+ plt.figure(figsize=(12, 8))
334
+ plt.title("Training Progress")
335
+ plt.plot(train_loss[:], label="Training Loss")
336
+ plt.plot(val_loss[:], label="Validation Loss")
337
+ plt.xlabel("Epoch")
338
+ plt.ylabel("Loss")
339
+ plt.legend()
340
+ plt.savefig(f"{odir}/results/__model_training.png")
341
+ plt.close()
342
+
343
+ plt.figure(figsize=(12, 8))
344
+ plt.title("Training Progress")
345
+ plt.plot(train_lcl_loss[25:], label="Training Loss")
346
+ plt.plot(val_lcl_loss[warmup:], label="Validation Loss")
347
+ plt.xlabel("Epoch")
348
+ plt.ylabel("Loss")
349
+ plt.legend()
350
+ plt.savefig(f"{odir}/results/__model_lcl_training.png")
351
+ plt.close()
352
+
353
+ plt.figure(figsize=(12, 8))
354
+ plt.title("Training Progress")
355
+ plt.plot(edge_values[25:])
356
+ plt.xlabel("Epoch")
357
+ plt.ylabel("EDGE Gradient")
358
+ plt.savefig(f"{odir}/results/__edge_training.png")
359
+ plt.close()
360
+
361
+
362
+ if __name__ == "__main__":
363
+
364
+ parser = argparse.ArgumentParser(description="Noise2Inverse with 2.5D")
365
+ parser.add_argument("-gpus", type=str, default="0", help="list of visiable GPUs")
366
+ parser.add_argument(
367
+ "-verbose", type=int, default=1, help="1:print to terminal; 0: redirect to file"
368
+ )
369
+ parser.add_argument("-config", type=str, required=True, help="path to config yaml file")
370
+
371
+ args, unparsed = parser.parse_known_args()
372
+
373
+ if len(unparsed) > 0:
374
+ print("Unrecognized argument(s): \n%s \nProgram exiting ... ..." % "\n".join(unparsed))
375
+ exit(0)
376
+
377
+ if len(args.gpus) > 0:
378
+ os.environ["CUDA_VISIBLE_DEVICES"] = args.gpus
379
+
380
+ logical_cpus = os.cpu_count()
381
+ os.environ["OMP_NUM_THREADS"] = str(logical_cpus)
382
+
383
+ logging.getLogger("matplotlib.font_manager").disabled = True
384
+
385
+ main(args)
@@ -0,0 +1,68 @@
1
+ """
2
+ utils.py — Small image-saving and command-line helper utilities.
3
+
4
+ Helpers shared across the project: write an array to a TIFF or 8-bit PNG preview
5
+ (`save2img`), save an RGB preview (`save2img_rgb`), rescale an array to uint8
6
+ (`scale2uint8`), and parse boolean command-line strings (`str2bool`).
7
+
8
+ Author: Cameron Renteria <crentb23@gmail.com>
9
+ License: Apache-2.0 (see LICENSE)
10
+ """
11
+
12
+ import argparse
13
+
14
+ import skimage.io
15
+ from matplotlib import pyplot as plt
16
+
17
+
18
+ def save2img_rgb(img_data, img_fn):
19
+ """Save an array as an RGB PNG preview, sized to the data on a black background."""
20
+ plt.figure(figsize=(img_data.shape[1] / 10.0, img_data.shape[0] / 10.0))
21
+ plt.axes([0, 0, 1, 1])
22
+ plt.imshow(
23
+ img_data,
24
+ )
25
+ plt.axis("off")
26
+ plt.savefig(img_fn, facecolor="black", edgecolor="black", dpi=10)
27
+ plt.close()
28
+
29
+
30
+ def save2img(d_img, fn):
31
+ """Save an array to disk.
32
+
33
+ If `fn` ends in 'tiff' the raw values are written unchanged; otherwise the data
34
+ is min-max scaled to 0-255 and saved as an 8-bit image (e.g. a PNG preview).
35
+ """
36
+ if fn[-4:] == "tiff":
37
+ img_norm = d_img.copy()
38
+ else:
39
+ _min, _max = d_img.min(), d_img.max()
40
+ if _max == _min:
41
+ img_norm = d_img - _max
42
+ else:
43
+ img_norm = (d_img - _min) * 255.0 / (_max - _min)
44
+ img_norm = img_norm.astype("uint8")
45
+ skimage.io.imsave(fn, img_norm, check_contrast=False)
46
+
47
+
48
+ def scale2uint8(_img):
49
+ """Min-max scale an array to the 0-255 uint8 range (a constant image maps to 0)."""
50
+ _min, _max = _img.min(), _img.max()
51
+ if _max == _min:
52
+ _img_s = _img - _max
53
+ else:
54
+ _img_s = (_img - _min) * 255.0 / (_max - _min)
55
+ _img_s = _img_s.astype("uint8")
56
+ return _img_s
57
+
58
+
59
+ def str2bool(v):
60
+ """Parse a truthy/falsy command-line string into a bool (for use with argparse)."""
61
+ if isinstance(v, bool):
62
+ return v
63
+ if v.lower() in ("yes", "true", "t", "y", "1"):
64
+ return True
65
+ elif v.lower() in ("no", "false", "f", "n", "0"):
66
+ return False
67
+ else:
68
+ raise argparse.ArgumentTypeError("Boolean value expected.")
ct_seg/labeling.py ADDED
@@ -0,0 +1,145 @@
1
+ """
2
+ Napari-based Labeling Tool for CT Segmentation
3
+ ===============================================
4
+ Opens a tiff stack in Napari with a labels layer for manual annotation.
5
+ Supports loading existing masks, painting labels, and saving.
6
+
7
+ Usage:
8
+ python label_tool.py --images /path/to/tiffs
9
+ python label_tool.py --images /path/to/tiffs --masks /path/to/existing_masks
10
+ python label_tool.py --images /path/to/tiffs --num_classes 4 --slice_range 500 600
11
+
12
+ Controls in Napari:
13
+ - Select the "Labels" layer in the layer list
14
+ - Use the paint brush (press B) to draw labels
15
+ - Use the eraser (press E) to erase
16
+ - Number keys (1, 2, 3, ...) select the label class
17
+ - Scroll through slices with the slider at the bottom
18
+ - Press Ctrl+S or use File > Save to save labels
19
+
20
+ When you close Napari, masks are automatically saved.
21
+
22
+ Requirements:
23
+ pip install napari[all]
24
+ """
25
+
26
+ import argparse
27
+ import re
28
+ import sys
29
+ from pathlib import Path
30
+
31
+ import numpy as np
32
+ import tifffile
33
+
34
+
35
+ def natural_sorted(paths):
36
+ def key(x):
37
+ return [int(c) if c.isdigit() else c for c in re.split(r"([0-9]+)", str(x))]
38
+
39
+ return sorted(paths, key=key)
40
+
41
+
42
+ def load_stack(dir_path, start=None, end=None):
43
+ dir_path = Path(dir_path)
44
+ paths = natural_sorted(list(dir_path.glob("*.tif*")))
45
+ if start is not None and end is not None:
46
+ paths = paths[start:end]
47
+ if not paths:
48
+ raise FileNotFoundError(f"No tiff files in {dir_path}")
49
+ from tqdm import tqdm
50
+
51
+ imgs = np.stack([tifffile.imread(str(p)) for p in tqdm(paths, desc="Loading")])
52
+ return imgs, paths
53
+
54
+
55
+ def save_masks(masks, out_dir, offset=0):
56
+ out_dir = Path(out_dir)
57
+ out_dir.mkdir(parents=True, exist_ok=True)
58
+ from tqdm import tqdm
59
+
60
+ for i in tqdm(range(masks.shape[0]), desc="Saving masks"):
61
+ tifffile.imwrite(str(out_dir / f"{i + offset:05d}.tiff"), masks[i].astype(np.uint8))
62
+ print(f"Saved {masks.shape[0]} masks to {out_dir}")
63
+
64
+
65
+ def main():
66
+ parser = argparse.ArgumentParser(description="Napari Labeling Tool for CT Segmentation")
67
+ parser.add_argument("--images", required=True, help="Directory of image tiffs")
68
+ parser.add_argument("--masks", default=None, help="Directory of existing mask tiffs (optional)")
69
+ parser.add_argument(
70
+ "--output", default=None, help="Output directory for masks (default: images_masks)"
71
+ )
72
+ parser.add_argument("--num_classes", type=int, default=4, help="Number of classes (default: 4)")
73
+ parser.add_argument("--slice_range", nargs=2, type=int, default=None, metavar=("START", "END"))
74
+ args = parser.parse_args()
75
+
76
+ try:
77
+ import napari
78
+ except ImportError:
79
+ print("Napari is required for the labeling tool.")
80
+ print("Install with: pip install 'napari[all]'")
81
+ sys.exit(1)
82
+
83
+ # Setup output
84
+ if args.output is None:
85
+ args.output = str(Path(args.images).parent / (Path(args.images).name + "_masks"))
86
+
87
+ # Load images
88
+ start = args.slice_range[0] if args.slice_range else None
89
+ end = args.slice_range[1] if args.slice_range else None
90
+ images, paths = load_stack(args.images, start=start, end=end)
91
+ print(f"Loaded {images.shape[0]} images of size {images.shape[1]}x{images.shape[2]}")
92
+
93
+ # Load or create masks
94
+ if args.masks:
95
+ masks, _ = load_stack(args.masks, start=start, end=end)
96
+ print(f"Loaded existing masks with classes: {np.unique(masks)}")
97
+ else:
98
+ masks = np.zeros(images.shape, dtype=np.uint8)
99
+ print("Created empty mask volume")
100
+
101
+ # Normalize images for display
102
+ p2, p98 = np.percentile(images, (2, 98))
103
+ display_imgs = np.clip(images, p2, p98)
104
+
105
+ # Color map for labels
106
+ label_colors = {
107
+ 0: [0, 0, 0, 0], # transparent background
108
+ 1: [1, 0.2, 0.2, 0.6], # red
109
+ 2: [0.2, 1, 0.2, 0.6], # green
110
+ 3: [0.2, 0.4, 1, 0.6], # blue
111
+ 4: [1, 1, 0.2, 0.6], # yellow
112
+ 5: [1, 0.2, 1, 0.6], # magenta
113
+ 6: [0.2, 1, 1, 0.6], # cyan
114
+ }
115
+
116
+ # Launch Napari
117
+ viewer = napari.Viewer(title="CT Segmentation Labeling Tool")
118
+ viewer.add_image(display_imgs, name="CT Image", colormap="gray")
119
+ labels_layer = viewer.add_labels(masks, name="Labels", color=label_colors)
120
+
121
+ # Set brush size
122
+ labels_layer.brush_size = 10
123
+
124
+ print(f"\n{'='*50}")
125
+ print("Napari Labeling Tool")
126
+ print(f"{'='*50}")
127
+ print(f"Classes: {args.num_classes}")
128
+ print("Controls:")
129
+ print(" B - Paint brush")
130
+ print(" E - Eraser")
131
+ print(" 1,2,3... - Select label class")
132
+ print(" Scroll - Navigate slices")
133
+ print(f"\nClose Napari to save masks to: {args.output}")
134
+ print(f"{'='*50}\n")
135
+
136
+ napari.run()
137
+
138
+ # Save on close
139
+ final_masks = labels_layer.data
140
+ offset = start if start else 0
141
+ save_masks(final_masks, args.output, offset=offset)
142
+
143
+
144
+ if __name__ == "__main__":
145
+ main()