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/denoise/train.py
ADDED
|
@@ -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)
|
ct_seg/denoise/utils.py
ADDED
|
@@ -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()
|