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
|
@@ -0,0 +1,164 @@
|
|
|
1
|
+
"""
|
|
2
|
+
denoise_volume.py — Denoise a full CT volume (or a slice range) with 2.5D N2I.
|
|
3
|
+
|
|
4
|
+
Loads the best edge-score checkpoint from training (main.py) and denoises every
|
|
5
|
+
slice of the full reconstruction, writing the result as a TIFF stack into a
|
|
6
|
+
`denoised_volume/` folder next to the reconstructions. A contiguous slice range
|
|
7
|
+
may be given to denoise only part of the volume (useful for quick evaluation); if
|
|
8
|
+
omitted, the whole volume is processed.
|
|
9
|
+
|
|
10
|
+
To use the GPU efficiently the script auto-tunes the inference batch size for the
|
|
11
|
+
available GPU memory, model size, and image dimensions (see
|
|
12
|
+
InferenceBatchSizeOptimizer in data.py). Normalization uses the mean/std that
|
|
13
|
+
training wrote back into the config file.
|
|
14
|
+
|
|
15
|
+
Inputs : a YAML config (with mean4norm/std4norm filled in by training) and an
|
|
16
|
+
optional [start_slice, end_slice] range.
|
|
17
|
+
Outputs : <reconstruction_dir>/denoised_volume/<index>.tiff stack
|
|
18
|
+
(this folder is deleted and recreated on every run).
|
|
19
|
+
|
|
20
|
+
Usage:
|
|
21
|
+
python denoise_volume.py -gpus=0 -config=/path/to/config.yaml -start_slice=500 -end_slice=600
|
|
22
|
+
(normally launched via denoise_volume.sh; omit the range to denoise everything)
|
|
23
|
+
|
|
24
|
+
Author: Cameron Renteria <crentb23@gmail.com>
|
|
25
|
+
License: Apache-2.0 (see LICENSE)
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
import argparse
|
|
29
|
+
import logging
|
|
30
|
+
import os
|
|
31
|
+
import shutil
|
|
32
|
+
import sys
|
|
33
|
+
import time
|
|
34
|
+
import warnings
|
|
35
|
+
|
|
36
|
+
import numpy as np
|
|
37
|
+
import torch
|
|
38
|
+
import yaml
|
|
39
|
+
from torch.utils.data import DataLoader
|
|
40
|
+
from tqdm import tqdm
|
|
41
|
+
|
|
42
|
+
from ct_seg.denoise import tiffs
|
|
43
|
+
from ct_seg.denoise.data import InferenceBatchSizeOptimizer, TomoDatasetInfer
|
|
44
|
+
from ct_seg.denoise.model import unet_ns_gn
|
|
45
|
+
|
|
46
|
+
warnings.filterwarnings("ignore")
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def main(args):
|
|
50
|
+
"""Load the trained model and denoise the configured volume (or slice range)."""
|
|
51
|
+
|
|
52
|
+
# Read the YAML file
|
|
53
|
+
with open(args.config, "r") as file:
|
|
54
|
+
params = yaml.safe_load(file)
|
|
55
|
+
|
|
56
|
+
# setup output directory
|
|
57
|
+
output_dir = params["dataset"]["directory_to_reconstructions"] + "/" "denoised_volume"
|
|
58
|
+
if os.path.isdir(output_dir):
|
|
59
|
+
shutil.rmtree(output_dir)
|
|
60
|
+
os.mkdir(output_dir)
|
|
61
|
+
|
|
62
|
+
# setup cuda device
|
|
63
|
+
dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
64
|
+
|
|
65
|
+
# load in model
|
|
66
|
+
path_to_mdl = (
|
|
67
|
+
params["dataset"]["directory_to_reconstructions"]
|
|
68
|
+
+ "/"
|
|
69
|
+
+ "TrainOutput"
|
|
70
|
+
+ "/"
|
|
71
|
+
+ "best_edge_model.pth"
|
|
72
|
+
)
|
|
73
|
+
# weights_only=True: restrict deserialization to tensors/containers.
|
|
74
|
+
# torch.load unpickles arbitrary Python by default, so a malicious
|
|
75
|
+
# checkpoint file would execute code on load (CWE-502 / bandit B614).
|
|
76
|
+
checkpoint = torch.load(path_to_mdl, map_location=torch.device("cpu"), weights_only=True)
|
|
77
|
+
model = unet_ns_gn(ich=5, start_filter_size=16, channels_per_group=8)
|
|
78
|
+
model.load_state_dict(checkpoint["model_state_dict"])
|
|
79
|
+
model.to(dev).eval()
|
|
80
|
+
|
|
81
|
+
print("\nLoading data into CPU memory, it will take a while ... ...")
|
|
82
|
+
|
|
83
|
+
# load in data
|
|
84
|
+
ds_test = TomoDatasetInfer(
|
|
85
|
+
params=params, start_slice=args.start_slice, end_slice=args.end_slice
|
|
86
|
+
)
|
|
87
|
+
print(
|
|
88
|
+
f"\nLoaded in {ds_test.reconstruction.shape[0]} slices of size {ds_test.reconstruction.shape[1]}x{ds_test.reconstruction.shape[2]}.\n"
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
# determine optimal batch size given GPU system memory, model size, and image size
|
|
92
|
+
optimal_batch_size = InferenceBatchSizeOptimizer(
|
|
93
|
+
model=model,
|
|
94
|
+
input_shape=ds_test.reconstruction[0].shape,
|
|
95
|
+
device=dev,
|
|
96
|
+
max_batch_size=16,
|
|
97
|
+
precision="fp32",
|
|
98
|
+
)
|
|
99
|
+
stats = optimal_batch_size.profile()
|
|
100
|
+
mbsz = stats["optimal_batch_size"]
|
|
101
|
+
|
|
102
|
+
dl_test = DataLoader(
|
|
103
|
+
dataset=ds_test,
|
|
104
|
+
batch_size=mbsz,
|
|
105
|
+
shuffle=False,
|
|
106
|
+
num_workers=4,
|
|
107
|
+
drop_last=False,
|
|
108
|
+
prefetch_factor=6,
|
|
109
|
+
pin_memory=True,
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
# initialize empty array for denoised volume
|
|
113
|
+
preds = np.zeros_like(dl_test.dataset.reconstruction)
|
|
114
|
+
insert_cnt = 0
|
|
115
|
+
# denoise volume
|
|
116
|
+
print("Processing data ...")
|
|
117
|
+
with torch.no_grad():
|
|
118
|
+
for X in tqdm(dl_test):
|
|
119
|
+
output = model(X.to(dev)).cpu().squeeze(dim=1).numpy()
|
|
120
|
+
|
|
121
|
+
preds[insert_cnt : (insert_cnt + X.shape[0])] = output
|
|
122
|
+
insert_cnt += X.shape[0]
|
|
123
|
+
|
|
124
|
+
# rescale volume
|
|
125
|
+
preds = preds * params["dataset"]["std4norm"] + params["dataset"]["mean4norm"]
|
|
126
|
+
|
|
127
|
+
# save volume
|
|
128
|
+
print("\nSaving data ...")
|
|
129
|
+
if len(args.start_slice) == 0:
|
|
130
|
+
tiffs.save_stack(output_dir, preds)
|
|
131
|
+
else:
|
|
132
|
+
# Save the processed sub volume with the right tiff number
|
|
133
|
+
tiffs.save_stack(output_dir, preds, offset=int(args.start_slice))
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
if __name__ == "__main__":
|
|
137
|
+
|
|
138
|
+
parser = argparse.ArgumentParser(description="Inference for 2.5D Noise2Inverse")
|
|
139
|
+
parser.add_argument("-gpus", type=str, default="0", help="list of visiable GPUs")
|
|
140
|
+
parser.add_argument("-start_slice", type=str, default=0, help="minibatch size")
|
|
141
|
+
parser.add_argument("-end_slice", type=str, default=None, help="minibatch size")
|
|
142
|
+
parser.add_argument("-config", type=str, required=True, help="path to config yaml file")
|
|
143
|
+
parser.add_argument(
|
|
144
|
+
"-verbose", type=int, default=1, help="1:print to terminal; 0: redirect to file"
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
args, unparsed = parser.parse_known_args()
|
|
148
|
+
|
|
149
|
+
if len(unparsed) > 0:
|
|
150
|
+
print("Unrecognized argument(s): \n%s \nProgram exiting ... ... " % "\n".join(unparsed))
|
|
151
|
+
exit(0)
|
|
152
|
+
|
|
153
|
+
if len(args.gpus) > 0:
|
|
154
|
+
os.environ["CUDA_VISIBLE_DEVICES"] = args.gpus
|
|
155
|
+
|
|
156
|
+
logging.getLogger("matplotlib.font_manager").disabled = True
|
|
157
|
+
logging.getLogger("matplotlib").setLevel(level=logging.CRITICAL)
|
|
158
|
+
if args.verbose:
|
|
159
|
+
logging.getLogger().addHandler(logging.StreamHandler(sys.stdout))
|
|
160
|
+
|
|
161
|
+
start_time = time.time()
|
|
162
|
+
main(args)
|
|
163
|
+
inference_time = time.time() - start_time
|
|
164
|
+
print(f"\nInference Time: {inference_time:.4f} seconds\n")
|
ct_seg/denoise/eval.py
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
"""
|
|
2
|
+
eval.py — Laplacian-based edge/sharpness score for monitoring denoising quality.
|
|
3
|
+
|
|
4
|
+
Provides `laplacian_score_batch`, the metric main.py uses during validation to
|
|
5
|
+
track how sharp the denoised output is. For each image it measures edge contrast
|
|
6
|
+
as the ratio of mean |Laplacian| in edge pixels to that in flat pixels; for very
|
|
7
|
+
flat images (low Laplacian-histogram entropy) it instead rewards smoothness. A
|
|
8
|
+
higher score means sharper, better-resolved edges, and main.py keeps the
|
|
9
|
+
checkpoint with the highest score (best_edge_model.pth).
|
|
10
|
+
|
|
11
|
+
Author: Cameron Renteria <crentb23@gmail.com>
|
|
12
|
+
License: Apache-2.0 (see LICENSE)
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
import numpy as np
|
|
16
|
+
import torch
|
|
17
|
+
import torch.nn.functional as F
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def laplacian_batch(x):
|
|
21
|
+
"""Apply a 3x3 discrete Laplacian (edge) filter to a batch of images.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
x: tensor [B, C, H, W]; the same kernel is applied per channel (grouped conv).
|
|
25
|
+
Returns:
|
|
26
|
+
The Laplacian response, same shape as `x`.
|
|
27
|
+
"""
|
|
28
|
+
# x: [B, C, H, W] (assumes grayscale: C=1)
|
|
29
|
+
kernel = torch.tensor([[0, 1, 0], [1, -4, 1], [0, 1, 0]], dtype=x.dtype, device=x.device).view(
|
|
30
|
+
1, 1, 3, 3
|
|
31
|
+
)
|
|
32
|
+
kernel = kernel.repeat(x.shape[1], 1, 1, 1) # for grouped conv
|
|
33
|
+
return F.conv2d(x, kernel, padding=1, groups=x.shape[1])
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def laplacian_entropy_map(lap, bins=256):
|
|
37
|
+
"""Per-image Shannon entropy of the Laplacian-magnitude histogram.
|
|
38
|
+
|
|
39
|
+
Higher entropy means a busier (more textured/edge-rich) image. Returns a 1-D
|
|
40
|
+
tensor of length B, one value per image in the batch.
|
|
41
|
+
"""
|
|
42
|
+
# Compute entropy for each image in batch
|
|
43
|
+
B = lap.shape[0]
|
|
44
|
+
entropies = []
|
|
45
|
+
for i in range(B):
|
|
46
|
+
hist = torch.histc(lap[i, 0], bins=bins, min=0, max=lap[i, 0].max())
|
|
47
|
+
hist = hist / hist.sum()
|
|
48
|
+
hist = hist + 1e-8 # avoid log(0)
|
|
49
|
+
entropy = -torch.sum(hist * torch.log(hist))
|
|
50
|
+
entropies.append(entropy.item())
|
|
51
|
+
return torch.tensor(entropies, device=lap.device)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def laplacian_score_batch(batch, entropy_thresh=0.2, q=0.9):
|
|
55
|
+
"""Mean edge-quality score over a batch (higher = sharper edges).
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
batch : tensor [B, 1, H, W] of images to score.
|
|
59
|
+
entropy_thresh: below this Laplacian-histogram entropy an image is treated
|
|
60
|
+
as "flat" and scored on smoothness instead of edge contrast.
|
|
61
|
+
q : quantile that separates edge pixels from flat pixels.
|
|
62
|
+
Returns:
|
|
63
|
+
float : the mean per-image score across the batch.
|
|
64
|
+
"""
|
|
65
|
+
# batch: [B, 1, H, W]
|
|
66
|
+
lap = torch.abs(laplacian_batch(batch)) # [B, 1, H, W]
|
|
67
|
+
B = lap.shape[0]
|
|
68
|
+
scores = []
|
|
69
|
+
|
|
70
|
+
entropies = laplacian_entropy_map(lap)
|
|
71
|
+
|
|
72
|
+
for i in range(B):
|
|
73
|
+
lap_i = lap[i, 0] # [H, W]
|
|
74
|
+
entropy = entropies[i]
|
|
75
|
+
|
|
76
|
+
if entropy < entropy_thresh:
|
|
77
|
+
# flat image: reward smoothness
|
|
78
|
+
# score = -torch.mean(lap_i).item()
|
|
79
|
+
smoothness = 1.0 / (lap_i.mean().item() + 1e-6)
|
|
80
|
+
scores.append(smoothness)
|
|
81
|
+
else:
|
|
82
|
+
# compute threshold using quantile
|
|
83
|
+
threshold = torch.quantile(lap_i, q)
|
|
84
|
+
edge_mask = lap_i > threshold
|
|
85
|
+
flat_mask = ~edge_mask
|
|
86
|
+
|
|
87
|
+
edge_score = lap_i[edge_mask].mean() if edge_mask.any() else 0.0
|
|
88
|
+
flat_score = lap_i[flat_mask].mean() if flat_mask.any() else 1e-6 # prevent div 0
|
|
89
|
+
|
|
90
|
+
contrast = edge_score / (flat_score + 1e-6)
|
|
91
|
+
scores.append(contrast)
|
|
92
|
+
|
|
93
|
+
return float(np.mean(scores))
|
|
94
|
+
# return np.array(scores).astype(np.float32)
|
ct_seg/denoise/loss.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
"""
|
|
2
|
+
loss.py — Laplacian Contrast Loss (LCL) for edge-aware CT denoising.
|
|
3
|
+
|
|
4
|
+
Defines the auxiliary loss used by main.py after the L1 warm-up. The Laplacian
|
|
5
|
+
highlights edges; LCL compares the mean Laplacian magnitude of "edge" pixels (the
|
|
6
|
+
top 20% by |Laplacian|) against "flat" pixels and returns flat/edge. Minimizing
|
|
7
|
+
this ratio pushes the network to keep edges sharp (high Laplacian) while smoothing
|
|
8
|
+
flat regions (low Laplacian), counteracting the over-smoothing that plain L1
|
|
9
|
+
denoising tends to produce.
|
|
10
|
+
|
|
11
|
+
Author: Cameron Renteria <crentb23@gmail.com>
|
|
12
|
+
License: Apache-2.0 (see LICENSE)
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
import torch
|
|
16
|
+
import torch.nn as nn
|
|
17
|
+
import torch.nn.functional as F
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def laplacian_batch(x):
|
|
21
|
+
"""Apply a 3x3 discrete Laplacian (edge) filter to a batch of images.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
x: tensor [B, C, H, W]; the same kernel is applied per channel (grouped conv).
|
|
25
|
+
Returns:
|
|
26
|
+
The Laplacian response, same shape as `x`.
|
|
27
|
+
"""
|
|
28
|
+
# x: [B, C, H, W] (assumes grayscale: C=1)
|
|
29
|
+
kernel = torch.tensor([[0, 1, 0], [1, -4, 1], [0, 1, 0]], dtype=x.dtype, device=x.device).view(
|
|
30
|
+
1, 1, 3, 3
|
|
31
|
+
)
|
|
32
|
+
kernel = kernel.repeat(x.shape[1], 1, 1, 1) # for grouped conv
|
|
33
|
+
return F.conv2d(x, kernel, padding=1, groups=x.shape[1])
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class LCL(nn.Module):
|
|
37
|
+
"""Laplacian Contrast Loss (see module docstring).
|
|
38
|
+
|
|
39
|
+
forward(pred) returns flat_mean / (edge_mean + eps), where edge pixels are the
|
|
40
|
+
top 20% by |Laplacian| and flat pixels are the rest. Lower is better: it favors
|
|
41
|
+
sharp edges relative to flat regions.
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
):
|
|
47
|
+
super(LCL, self).__init__()
|
|
48
|
+
|
|
49
|
+
def forward(self, pred):
|
|
50
|
+
L = torch.abs(laplacian_batch(pred))
|
|
51
|
+
threshold = torch.quantile(L, 0.80)
|
|
52
|
+
# Split pixels into 'edge' (top 20% by Laplacian magnitude) and 'flat' (the rest).
|
|
53
|
+
edge_mask = L > threshold
|
|
54
|
+
flat_mask = ~edge_mask
|
|
55
|
+
|
|
56
|
+
edge_mean = L[edge_mask].mean() if edge_mask.any() else 0.0
|
|
57
|
+
flat_mean = L[flat_mask].mean() if flat_mask.any() else 1e-6
|
|
58
|
+
|
|
59
|
+
# Encourage this ratio to grow (i.e., minimize the inverse)
|
|
60
|
+
return flat_mean / (edge_mean + 1e-6)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def laplacian_entropy_map(lap, bins=256):
|
|
64
|
+
"""Per-image Shannon entropy of the Laplacian-magnitude histogram.
|
|
65
|
+
|
|
66
|
+
Higher entropy means a busier (more textured/edge-rich) image. Returns a 1-D
|
|
67
|
+
tensor of length B, one value per image in the batch.
|
|
68
|
+
"""
|
|
69
|
+
# Compute entropy for each image in batch
|
|
70
|
+
B = lap.shape[0]
|
|
71
|
+
entropies = []
|
|
72
|
+
for i in range(B):
|
|
73
|
+
hist = torch.histc(lap[i, 0], bins=bins, min=0, max=lap[i, 0].max().item())
|
|
74
|
+
hist = hist / hist.sum()
|
|
75
|
+
hist = hist + 1e-8 # avoid log(0)
|
|
76
|
+
entropy = -torch.sum(hist * torch.log(hist))
|
|
77
|
+
entropies.append(entropy.item())
|
|
78
|
+
return torch.tensor(entropies, device=lap.device)
|
ct_seg/denoise/model.py
ADDED
|
@@ -0,0 +1,205 @@
|
|
|
1
|
+
"""
|
|
2
|
+
model.py — U-Net architecture for 2.5D Noise2Inverse denoising.
|
|
3
|
+
|
|
4
|
+
Defines a compact 2D U-Net used as the denoiser. The variant used by the project
|
|
5
|
+
is `unet_ns_gn` ("no-skip, group-norm"): a U-Net WITHOUT the usual skip
|
|
6
|
+
connections between encoder and decoder, using Group Normalization and LeakyReLU
|
|
7
|
+
throughout. Empirically this no-skip design is a robust denoiser across different
|
|
8
|
+
CT samples (skip connections tend to let high-frequency noise bypass the network).
|
|
9
|
+
|
|
10
|
+
The network takes a stack of adjacent slices as input channels (5 for the 2.5D
|
|
11
|
+
setup) and outputs a single denoised slice. Helper modules:
|
|
12
|
+
unet_box_gn - double 3x3 conv block (conv -> GroupNorm -> LeakyReLU, x2)
|
|
13
|
+
unet_bottleneck_gn - single 3x3 conv block at the bottleneck
|
|
14
|
+
unet_down - 2x2 max-pool downsampling
|
|
15
|
+
unet_up - nearest-neighbour 2x upsampling
|
|
16
|
+
unet_ns_gn - the full encoder/bottleneck/decoder network
|
|
17
|
+
|
|
18
|
+
Author: Cameron Renteria <crentb23@gmail.com>
|
|
19
|
+
License: Apache-2.0 (see LICENSE)
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
import torch
|
|
23
|
+
import torch.nn as nn
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class unet_box_gn(torch.nn.Module):
|
|
27
|
+
"""Double convolution block: (Conv3x3 -> GroupNorm -> LeakyReLU) applied twice.
|
|
28
|
+
|
|
29
|
+
Preserves spatial size (padding=1) while mapping `in_ch` -> `out_ch` channels.
|
|
30
|
+
`groups` sets the number of GroupNorm groups.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
def __init__(self, in_ch, out_ch, groups):
|
|
34
|
+
super().__init__()
|
|
35
|
+
self.double_conv = torch.nn.Sequential(
|
|
36
|
+
torch.nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
|
|
37
|
+
nn.GroupNorm(num_groups=groups, num_channels=out_ch),
|
|
38
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
39
|
+
torch.nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),
|
|
40
|
+
nn.GroupNorm(num_groups=groups, num_channels=out_ch),
|
|
41
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
def forward(self, x):
|
|
45
|
+
return self.double_conv(x)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class unet_bottleneck_gn(torch.nn.Module):
|
|
49
|
+
"""Single convolution block (Conv3x3 -> GroupNorm -> LeakyReLU) for the bottleneck."""
|
|
50
|
+
|
|
51
|
+
def __init__(self, in_ch, out_ch, groups):
|
|
52
|
+
super().__init__()
|
|
53
|
+
self.bn_conv = torch.nn.Sequential(
|
|
54
|
+
torch.nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
|
|
55
|
+
nn.GroupNorm(num_groups=groups, num_channels=out_ch),
|
|
56
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
def forward(self, x):
|
|
60
|
+
return self.bn_conv(x)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class unet_up(torch.nn.Module):
|
|
64
|
+
"""Upsampling block: nearest-neighbour 2x upsampling (no learned parameters).
|
|
65
|
+
|
|
66
|
+
The `ch` argument is accepted for interface symmetry but is unused.
|
|
67
|
+
"""
|
|
68
|
+
|
|
69
|
+
def __init__(
|
|
70
|
+
self,
|
|
71
|
+
ch,
|
|
72
|
+
):
|
|
73
|
+
super().__init__()
|
|
74
|
+
self.down_scale = torch.nn.Sequential(torch.nn.Upsample(scale_factor=2, mode="nearest"))
|
|
75
|
+
|
|
76
|
+
def forward(self, x):
|
|
77
|
+
return self.down_scale(x)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
class unet_down(torch.nn.Module):
|
|
81
|
+
"""Downsampling block: 2x2 max pooling (halves the spatial resolution)."""
|
|
82
|
+
|
|
83
|
+
def __init__(self, ch):
|
|
84
|
+
super().__init__()
|
|
85
|
+
self.maxpool = torch.nn.Sequential(
|
|
86
|
+
torch.nn.MaxPool2d(2),
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
def forward(self, x):
|
|
90
|
+
return self.maxpool(x)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
class unet_ns_gn(torch.nn.Module):
|
|
94
|
+
"""No-skip U-Net with Group Normalization (the project's denoiser).
|
|
95
|
+
|
|
96
|
+
Args:
|
|
97
|
+
start_filter_size : base channel width; deeper stages use multiples of it.
|
|
98
|
+
ich : number of input channels (5 for 2.5D: five adjacent slices).
|
|
99
|
+
och : number of output channels (1 = a single denoised slice).
|
|
100
|
+
channels_per_group: channels per GroupNorm group (sets num_groups internally).
|
|
101
|
+
|
|
102
|
+
The forward pass is a straight encoder -> bottleneck -> decoder with NO skip
|
|
103
|
+
connections; spatial resolution is restored purely by the upsampling stages.
|
|
104
|
+
"""
|
|
105
|
+
|
|
106
|
+
def __init__(self, start_filter_size, ich=1, och=1, channels_per_group=8):
|
|
107
|
+
super().__init__()
|
|
108
|
+
# Stem: 1x1 conv lifting the `ich` input channels up to start_filter_size.
|
|
109
|
+
self.in_box = torch.nn.Sequential(
|
|
110
|
+
torch.nn.Conv2d(ich, start_filter_size, kernel_size=1, padding=0),
|
|
111
|
+
nn.GroupNorm(
|
|
112
|
+
num_groups=int((start_filter_size) / channels_per_group),
|
|
113
|
+
num_channels=start_filter_size,
|
|
114
|
+
),
|
|
115
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
116
|
+
)
|
|
117
|
+
# --- Encoder: progressively widen channels and halve spatial size ---
|
|
118
|
+
self.box1 = unet_box_gn(
|
|
119
|
+
start_filter_size,
|
|
120
|
+
start_filter_size * 4,
|
|
121
|
+
groups=int((start_filter_size * 4) / channels_per_group),
|
|
122
|
+
)
|
|
123
|
+
self.down1 = unet_down(start_filter_size * 4)
|
|
124
|
+
|
|
125
|
+
self.box2 = unet_box_gn(
|
|
126
|
+
start_filter_size * 4,
|
|
127
|
+
start_filter_size * 8,
|
|
128
|
+
groups=int((start_filter_size * 8) / channels_per_group),
|
|
129
|
+
)
|
|
130
|
+
self.down2 = unet_down(start_filter_size * 8)
|
|
131
|
+
|
|
132
|
+
self.box3 = unet_box_gn(
|
|
133
|
+
start_filter_size * 8,
|
|
134
|
+
start_filter_size * 16,
|
|
135
|
+
groups=int((start_filter_size * 16) / channels_per_group),
|
|
136
|
+
)
|
|
137
|
+
self.down3 = unet_down(start_filter_size * 16)
|
|
138
|
+
|
|
139
|
+
# --- Bottleneck (lowest resolution) ---
|
|
140
|
+
self.bottleneck = unet_bottleneck_gn(
|
|
141
|
+
start_filter_size * 16,
|
|
142
|
+
start_filter_size * 16,
|
|
143
|
+
groups=int((start_filter_size * 16) / channels_per_group),
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
# --- Decoder: upsample back to full resolution (NO skip connections) ---
|
|
147
|
+
self.up1 = unet_up(start_filter_size * 16)
|
|
148
|
+
self.box4 = unet_box_gn(
|
|
149
|
+
start_filter_size * 16,
|
|
150
|
+
start_filter_size * 8,
|
|
151
|
+
groups=int((start_filter_size * 8) / channels_per_group),
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
self.up2 = unet_up(start_filter_size * 8)
|
|
155
|
+
self.box5 = unet_box_gn(
|
|
156
|
+
start_filter_size * 8,
|
|
157
|
+
start_filter_size * 4,
|
|
158
|
+
groups=int((start_filter_size * 4) / channels_per_group),
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
self.up3 = unet_up(start_filter_size * 4)
|
|
162
|
+
self.box6 = unet_box_gn(
|
|
163
|
+
start_filter_size * 4,
|
|
164
|
+
start_filter_size * 4,
|
|
165
|
+
groups=int((start_filter_size * 4) / channels_per_group),
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
# Output head: project down to `och` output channel(s).
|
|
169
|
+
self.out_layer = torch.nn.Sequential(
|
|
170
|
+
torch.nn.Conv2d(start_filter_size * 4, start_filter_size * 2, kernel_size=1, padding=0),
|
|
171
|
+
nn.GroupNorm(
|
|
172
|
+
num_groups=int((start_filter_size * 2) / channels_per_group),
|
|
173
|
+
num_channels=start_filter_size * 2,
|
|
174
|
+
),
|
|
175
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
176
|
+
torch.nn.Conv2d(start_filter_size * 2, och, kernel_size=1, padding=0),
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
def forward(self, x):
|
|
180
|
+
output = self.in_box(x)
|
|
181
|
+
|
|
182
|
+
output = self.box1(output)
|
|
183
|
+
output = self.down1(output)
|
|
184
|
+
|
|
185
|
+
output = self.box2(output)
|
|
186
|
+
output = self.down2(output)
|
|
187
|
+
|
|
188
|
+
output = self.box3(output)
|
|
189
|
+
output = self.down3(output)
|
|
190
|
+
|
|
191
|
+
output = self.bottleneck(output)
|
|
192
|
+
|
|
193
|
+
output = self.up1(output)
|
|
194
|
+
|
|
195
|
+
output = self.box4(output)
|
|
196
|
+
output = self.up2(output)
|
|
197
|
+
|
|
198
|
+
output = self.box5(output)
|
|
199
|
+
output = self.up3(output)
|
|
200
|
+
|
|
201
|
+
output = self.box6(output)
|
|
202
|
+
|
|
203
|
+
output = self.out_layer(output)
|
|
204
|
+
|
|
205
|
+
return output
|
ct_seg/denoise/tiffs.py
ADDED
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
"""
|
|
2
|
+
tiffs.py — TIFF stack input/output helpers for CT volumes.
|
|
3
|
+
|
|
4
|
+
Utilities for reading and writing the TIFF image stacks that make up a CT
|
|
5
|
+
reconstruction: natural (human) sorting of filenames, loading a directory of TIFFs
|
|
6
|
+
into a NumPy volume, globbing a directory for TIFFs, saving a volume back out as
|
|
7
|
+
numbered TIFFs, and loading a stack as a sinogram.
|
|
8
|
+
|
|
9
|
+
Author: Cameron Renteria <crentb23@gmail.com>
|
|
10
|
+
License: Apache-2.0 (see LICENSE)
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
import re
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
|
|
16
|
+
import numpy as np
|
|
17
|
+
import tifffile
|
|
18
|
+
from tqdm import tqdm
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def natural_sorted(paths):
|
|
22
|
+
"""Sort paths/strings in natural (human) order so that e.g. img2 precedes img10."""
|
|
23
|
+
|
|
24
|
+
def key(x):
|
|
25
|
+
return [int(c) if c.isdigit() else c for c in re.split("([0-9]+)", str(x))]
|
|
26
|
+
|
|
27
|
+
return sorted(paths, key=key)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# We use the following function to load a stack of images:
|
|
31
|
+
def load_stack(paths, binning=1, use_tqdm=True):
|
|
32
|
+
"""Load a stack of tiff files.
|
|
33
|
+
|
|
34
|
+
:param paths: paths to tiff files
|
|
35
|
+
:param binning: whether angles and projection images should be binned.
|
|
36
|
+
:returns: an np.array containing the values in the tiff files
|
|
37
|
+
:rtype: np.array
|
|
38
|
+
|
|
39
|
+
"""
|
|
40
|
+
# Read first image for shape and dtype information
|
|
41
|
+
paths = list(paths)
|
|
42
|
+
|
|
43
|
+
img0 = tifffile.imread(str(paths[0]))
|
|
44
|
+
img0 = img0[::binning, ::binning]
|
|
45
|
+
dtype = img0.dtype
|
|
46
|
+
# Create empty numpy array to hold result
|
|
47
|
+
imgs = np.empty((len(paths), *img0.shape), dtype=dtype)
|
|
48
|
+
|
|
49
|
+
for i in tqdm(range(len(paths))):
|
|
50
|
+
imgs[i] = tifffile.imread(str(paths[i]))[::binning, ::binning]
|
|
51
|
+
return imgs
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def glob(dir_path):
|
|
55
|
+
"""Expand path to list of all tiffs in directory
|
|
56
|
+
|
|
57
|
+
:param dir_path: directory
|
|
58
|
+
:returns:
|
|
59
|
+
:rtype:
|
|
60
|
+
|
|
61
|
+
"""
|
|
62
|
+
dir_path = Path(dir_path).expanduser().resolve()
|
|
63
|
+
return natural_sorted(dir_path.glob("*.tif*"))
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def save_stack(path, stack, offset=0, exist_ok=True, parents=False):
|
|
67
|
+
"""Write a 3-D volume to `path` as zero-padded numbered TIFFs (NNNNN.tiff).
|
|
68
|
+
|
|
69
|
+
`offset` shifts the starting file number so a denoised sub-range keeps the same
|
|
70
|
+
indices as the original slices.
|
|
71
|
+
"""
|
|
72
|
+
path = Path(path).expanduser().resolve()
|
|
73
|
+
path.mkdir(exist_ok=exist_ok, parents=parents)
|
|
74
|
+
for i in tqdm(range(stack.shape[0])):
|
|
75
|
+
opath = path / f"{i+offset:05d}.tiff"
|
|
76
|
+
tifffile.imwrite(str(opath), stack[i])
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def load_sino(paths, binning=1, dtype=None, flip_y=False):
|
|
80
|
+
"""Load a stack of tiff files into a sinogram
|
|
81
|
+
|
|
82
|
+
:param paths: paths to tiff files
|
|
83
|
+
:param binning: whether angles and projection images should be binned.
|
|
84
|
+
:returns: an np.array containing the values in the tiff files
|
|
85
|
+
:rtype: np.array
|
|
86
|
+
|
|
87
|
+
"""
|
|
88
|
+
# Read first image for shape and dtype information
|
|
89
|
+
paths = list(paths)
|
|
90
|
+
# print(paths[0])
|
|
91
|
+
# sys.exit()
|
|
92
|
+
img0 = tifffile.imread(str(paths[0]))
|
|
93
|
+
img0 = img0[::binning, ::binning]
|
|
94
|
+
if dtype is None:
|
|
95
|
+
dtype = img0.dtype
|
|
96
|
+
# Create empty numpy array to hold result
|
|
97
|
+
imgs = np.empty((img0.shape[0], len(paths), img0.shape[1]), dtype=dtype)
|
|
98
|
+
for i, p in tqdm(enumerate(paths)):
|
|
99
|
+
# Angles in the middle, "up" in front, "right" at the back.
|
|
100
|
+
if flip_y:
|
|
101
|
+
# Flip in the vertical direction
|
|
102
|
+
imgs[:, i, :] = tifffile.imread(str(p))[::-binning, ::binning]
|
|
103
|
+
else:
|
|
104
|
+
imgs[:, i, :] = tifffile.imread(str(p))[::binning, ::binning]
|
|
105
|
+
return imgs
|