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/model.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
"""
|
|
2
|
+
U-Net with Skip Connections for Semantic Segmentation of CT Images
|
|
3
|
+
==================================================================
|
|
4
|
+
Standard U-Net architecture with skip connections (concatenation),
|
|
5
|
+
Group Normalization, and configurable number of output classes.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
import torch.nn as nn
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class DoubleConv(nn.Module):
|
|
13
|
+
"""Conv3x3 -> GN -> LeakyReLU -> Conv3x3 -> GN -> LeakyReLU"""
|
|
14
|
+
|
|
15
|
+
def __init__(self, in_ch, out_ch, groups=8):
|
|
16
|
+
super().__init__()
|
|
17
|
+
# Ensure groups doesn't exceed channel count
|
|
18
|
+
g = min(groups, out_ch)
|
|
19
|
+
self.block = nn.Sequential(
|
|
20
|
+
nn.Conv2d(in_ch, out_ch, 3, padding=1),
|
|
21
|
+
nn.GroupNorm(num_groups=g, num_channels=out_ch),
|
|
22
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
23
|
+
nn.Conv2d(out_ch, out_ch, 3, padding=1),
|
|
24
|
+
nn.GroupNorm(num_groups=g, num_channels=out_ch),
|
|
25
|
+
nn.LeakyReLU(0.1, inplace=True),
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
def forward(self, x):
|
|
29
|
+
return self.block(x)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class UNetSegmentation(nn.Module):
|
|
33
|
+
"""
|
|
34
|
+
U-Net with skip connections for multi-class segmentation.
|
|
35
|
+
|
|
36
|
+
Parameters
|
|
37
|
+
----------
|
|
38
|
+
in_channels : int
|
|
39
|
+
Number of input channels (1 for grayscale, 5 for 2.5D).
|
|
40
|
+
num_classes : int
|
|
41
|
+
Number of segmentation classes.
|
|
42
|
+
base_filters : int
|
|
43
|
+
Number of filters in the first encoder level. Doubles at each level.
|
|
44
|
+
channels_per_group : int
|
|
45
|
+
Channels per group for Group Normalization.
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
def __init__(self, in_channels=1, num_classes=4, base_filters=32, channels_per_group=8):
|
|
49
|
+
super().__init__()
|
|
50
|
+
f = base_filters
|
|
51
|
+
|
|
52
|
+
# Encoder
|
|
53
|
+
self.enc1 = DoubleConv(in_channels, f, groups=min(channels_per_group, f))
|
|
54
|
+
self.enc2 = DoubleConv(f, f * 2, groups=min(channels_per_group, f * 2))
|
|
55
|
+
self.enc3 = DoubleConv(f * 2, f * 4, groups=min(channels_per_group, f * 4))
|
|
56
|
+
self.enc4 = DoubleConv(f * 4, f * 8, groups=min(channels_per_group, f * 8))
|
|
57
|
+
|
|
58
|
+
self.pool = nn.MaxPool2d(2)
|
|
59
|
+
|
|
60
|
+
# Bottleneck
|
|
61
|
+
self.bottleneck = DoubleConv(f * 8, f * 16, groups=min(channels_per_group, f * 16))
|
|
62
|
+
|
|
63
|
+
# Decoder (input channels = skip + upsampled)
|
|
64
|
+
self.up4 = nn.Upsample(scale_factor=2, mode="nearest")
|
|
65
|
+
self.dec4 = DoubleConv(f * 16 + f * 8, f * 8, groups=min(channels_per_group, f * 8))
|
|
66
|
+
|
|
67
|
+
self.up3 = nn.Upsample(scale_factor=2, mode="nearest")
|
|
68
|
+
self.dec3 = DoubleConv(f * 8 + f * 4, f * 4, groups=min(channels_per_group, f * 4))
|
|
69
|
+
|
|
70
|
+
self.up2 = nn.Upsample(scale_factor=2, mode="nearest")
|
|
71
|
+
self.dec2 = DoubleConv(f * 4 + f * 2, f * 2, groups=min(channels_per_group, f * 2))
|
|
72
|
+
|
|
73
|
+
self.up1 = nn.Upsample(scale_factor=2, mode="nearest")
|
|
74
|
+
self.dec1 = DoubleConv(f * 2 + f, f, groups=min(channels_per_group, f))
|
|
75
|
+
|
|
76
|
+
# Output
|
|
77
|
+
self.out_conv = nn.Conv2d(f, num_classes, kernel_size=1)
|
|
78
|
+
|
|
79
|
+
def forward(self, x):
|
|
80
|
+
# Encoder
|
|
81
|
+
e1 = self.enc1(x)
|
|
82
|
+
e2 = self.enc2(self.pool(e1))
|
|
83
|
+
e3 = self.enc3(self.pool(e2))
|
|
84
|
+
e4 = self.enc4(self.pool(e3))
|
|
85
|
+
|
|
86
|
+
# Bottleneck
|
|
87
|
+
b = self.bottleneck(self.pool(e4))
|
|
88
|
+
|
|
89
|
+
# Decoder with skip connections
|
|
90
|
+
d4 = self.dec4(torch.cat([self.up4(b), e4], dim=1))
|
|
91
|
+
d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1))
|
|
92
|
+
d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1))
|
|
93
|
+
d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1))
|
|
94
|
+
|
|
95
|
+
return self.out_conv(d1)
|
|
96
|
+
|
|
97
|
+
def predict(self, x):
|
|
98
|
+
"""Run forward pass and return class predictions (argmax)."""
|
|
99
|
+
with torch.no_grad():
|
|
100
|
+
logits = self.forward(x)
|
|
101
|
+
return torch.argmax(logits, dim=1)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def count_parameters(model):
|
|
105
|
+
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
if __name__ == "__main__":
|
|
109
|
+
# Quick test
|
|
110
|
+
model = UNetSegmentation(in_channels=1, num_classes=4, base_filters=32)
|
|
111
|
+
print(f"Parameters: {count_parameters(model):,}")
|
|
112
|
+
x = torch.randn(1, 1, 256, 256)
|
|
113
|
+
out = model(x)
|
|
114
|
+
print(f"Input: {x.shape}")
|
|
115
|
+
print(f"Output: {out.shape}") # [1, 4, 256, 256]
|
|
116
|
+
pred = model.predict(x)
|
|
117
|
+
print(f"Prediction: {pred.shape}") # [1, 256, 256]
|
ct_seg/segment.py
ADDED
|
@@ -0,0 +1,461 @@
|
|
|
1
|
+
"""
|
|
2
|
+
CT Image Stack Segmentation
|
|
3
|
+
============================
|
|
4
|
+
Supports both unsupervised and supervised segmentation of tiff stacks.
|
|
5
|
+
|
|
6
|
+
Unsupervised methods:
|
|
7
|
+
python segment.py --input /path/to/tiffs --method otsu --num_classes 4
|
|
8
|
+
python segment.py --input /path/to/tiffs --method kmeans --num_classes 3
|
|
9
|
+
python segment.py --input /path/to/tiffs --method gmm --num_classes 4
|
|
10
|
+
|
|
11
|
+
Supervised (U-Net):
|
|
12
|
+
python segment.py --input /path/to/tiffs --method unet --model /path/to/model.pth --num_classes 4
|
|
13
|
+
|
|
14
|
+
Options:
|
|
15
|
+
--slice_range 500 600 Process only slices 500-600
|
|
16
|
+
--enhance Apply contrast enhancement before segmentation
|
|
17
|
+
--save_overlay Save overlay images alongside masks
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
import argparse
|
|
21
|
+
import warnings
|
|
22
|
+
from pathlib import Path
|
|
23
|
+
|
|
24
|
+
import numpy as np
|
|
25
|
+
import tifffile
|
|
26
|
+
from tqdm import tqdm
|
|
27
|
+
|
|
28
|
+
warnings.filterwarnings("ignore")
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
# =============================================================================
|
|
32
|
+
# Tiff I/O (standalone, no dependency on N2I tiffs.py)
|
|
33
|
+
# =============================================================================
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def natural_sorted(paths):
|
|
37
|
+
import re
|
|
38
|
+
|
|
39
|
+
def key(x):
|
|
40
|
+
return [int(c) if c.isdigit() else c for c in re.split(r"([0-9]+)", str(x))]
|
|
41
|
+
|
|
42
|
+
return sorted(paths, key=key)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def load_image_stack(dir_path, start=None, end=None):
|
|
46
|
+
"""Load a directory of image files (tiff, png, jpg) into a 3D grayscale numpy array."""
|
|
47
|
+
from PIL import Image
|
|
48
|
+
|
|
49
|
+
dir_path = Path(dir_path)
|
|
50
|
+
# Support tiff, png, jpg
|
|
51
|
+
paths = []
|
|
52
|
+
for ext in ["*.tif*", "*.png", "*.jpg", "*.jpeg"]:
|
|
53
|
+
paths.extend(dir_path.glob(ext))
|
|
54
|
+
paths = natural_sorted(list(set(paths)))
|
|
55
|
+
|
|
56
|
+
if start is not None and end is not None:
|
|
57
|
+
paths = paths[start:end]
|
|
58
|
+
if len(paths) == 0:
|
|
59
|
+
raise FileNotFoundError(f"No image files found in {dir_path}")
|
|
60
|
+
|
|
61
|
+
# Read first image to determine shape
|
|
62
|
+
img0 = np.array(Image.open(str(paths[0])).convert("L")) # convert to grayscale
|
|
63
|
+
stack = np.empty((len(paths), *img0.shape), dtype=img0.dtype)
|
|
64
|
+
for i, p in enumerate(tqdm(paths, desc="Loading images")):
|
|
65
|
+
stack[i] = np.array(Image.open(str(p)).convert("L"))
|
|
66
|
+
return stack, paths
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def save_tiff_stack(stack, out_dir, num_classes, offset=0):
|
|
70
|
+
"""Save a 3D label array as RGB tiff files with vivid class colors."""
|
|
71
|
+
from PIL import Image
|
|
72
|
+
|
|
73
|
+
out_dir = Path(out_dir)
|
|
74
|
+
out_dir.mkdir(parents=True, exist_ok=True)
|
|
75
|
+
|
|
76
|
+
# Vivid colors per class
|
|
77
|
+
color_map = [
|
|
78
|
+
[0, 0, 0], # class 0: black
|
|
79
|
+
[255, 0, 0], # class 1: red
|
|
80
|
+
[0, 255, 0], # class 2: green
|
|
81
|
+
[0, 100, 255], # class 3: blue
|
|
82
|
+
[255, 255, 0], # class 4: yellow
|
|
83
|
+
[255, 0, 255], # class 5: magenta
|
|
84
|
+
[0, 255, 255], # class 6: cyan
|
|
85
|
+
[255, 165, 0], # class 7: orange
|
|
86
|
+
]
|
|
87
|
+
|
|
88
|
+
for i in tqdm(range(stack.shape[0]), desc="Saving colored masks"):
|
|
89
|
+
rgb = np.zeros((*stack[i].shape, 3), dtype=np.uint8)
|
|
90
|
+
for c in range(min(num_classes, len(color_map))):
|
|
91
|
+
rgb[stack[i] == c] = color_map[c]
|
|
92
|
+
Image.fromarray(rgb).save(str(out_dir / f"{i + offset:05d}.png"))
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
# =============================================================================
|
|
96
|
+
# Preprocessing
|
|
97
|
+
# =============================================================================
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def normalize_stack(stack):
|
|
101
|
+
"""Normalize stack to [0, 1] range using global min/max."""
|
|
102
|
+
smin, smax = stack.min(), stack.max()
|
|
103
|
+
if smax == smin:
|
|
104
|
+
return np.zeros_like(stack, dtype=np.float32)
|
|
105
|
+
return ((stack - smin) / (smax - smin)).astype(np.float32)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def enhance_contrast(stack, q_low=2, q_high=98):
|
|
109
|
+
"""Clip to percentile range and rescale to [0, 1]."""
|
|
110
|
+
p_low = np.percentile(stack, q_low)
|
|
111
|
+
p_high = np.percentile(stack, q_high)
|
|
112
|
+
clipped = np.clip(stack, p_low, p_high)
|
|
113
|
+
return normalize_stack(clipped)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
# =============================================================================
|
|
117
|
+
# Unsupervised Segmentation Methods
|
|
118
|
+
# =============================================================================
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def segment_otsu(stack, num_classes):
|
|
122
|
+
"""
|
|
123
|
+
Multi-Otsu thresholding.
|
|
124
|
+
Finds (num_classes - 1) thresholds that minimize intra-class variance.
|
|
125
|
+
"""
|
|
126
|
+
from skimage.filters import threshold_multiotsu
|
|
127
|
+
|
|
128
|
+
print(f"\nRunning Multi-Otsu segmentation with {num_classes} classes...")
|
|
129
|
+
|
|
130
|
+
# Compute thresholds on a subsample for speed
|
|
131
|
+
subsample = stack[:: max(1, len(stack) // 20)].ravel()
|
|
132
|
+
subsample = subsample[:: max(1, len(subsample) // 5_000_000)]
|
|
133
|
+
|
|
134
|
+
thresholds = threshold_multiotsu(subsample, classes=num_classes)
|
|
135
|
+
print(f"Thresholds: {thresholds}")
|
|
136
|
+
|
|
137
|
+
labels = np.digitize(stack, bins=thresholds).astype(np.uint8)
|
|
138
|
+
return labels, {"thresholds": thresholds.tolist()}
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def segment_kmeans(stack, num_classes):
|
|
142
|
+
"""
|
|
143
|
+
K-Means clustering on pixel intensities.
|
|
144
|
+
Clusters are sorted by centroid value so class 0 = darkest.
|
|
145
|
+
"""
|
|
146
|
+
from sklearn.cluster import MiniBatchKMeans
|
|
147
|
+
|
|
148
|
+
print(f"\nRunning K-Means segmentation with {num_classes} clusters...")
|
|
149
|
+
|
|
150
|
+
# Flatten and subsample for fitting
|
|
151
|
+
flat = stack.ravel().astype(np.float32)
|
|
152
|
+
subsample_idx = np.random.choice(len(flat), size=min(2_000_000, len(flat)), replace=False)
|
|
153
|
+
subsample = flat[subsample_idx].reshape(-1, 1)
|
|
154
|
+
|
|
155
|
+
kmeans = MiniBatchKMeans(n_clusters=num_classes, batch_size=10000, random_state=42, n_init=3)
|
|
156
|
+
kmeans.fit(subsample)
|
|
157
|
+
|
|
158
|
+
# Sort clusters by centroid value (class 0 = darkest)
|
|
159
|
+
sorted_idx = np.argsort(kmeans.cluster_centers_.ravel())
|
|
160
|
+
label_map = np.zeros(num_classes, dtype=np.uint8)
|
|
161
|
+
for new_label, old_label in enumerate(sorted_idx):
|
|
162
|
+
label_map[old_label] = new_label
|
|
163
|
+
|
|
164
|
+
# Predict all pixels
|
|
165
|
+
print("Assigning labels to full volume...")
|
|
166
|
+
labels = np.empty(stack.shape, dtype=np.uint8)
|
|
167
|
+
for i in tqdm(range(stack.shape[0]), desc="Segmenting slices"):
|
|
168
|
+
sl = stack[i].ravel().astype(np.float32).reshape(-1, 1)
|
|
169
|
+
pred = kmeans.predict(sl)
|
|
170
|
+
labels[i] = label_map[pred].reshape(stack[i].shape)
|
|
171
|
+
|
|
172
|
+
centroids = kmeans.cluster_centers_.ravel()[sorted_idx]
|
|
173
|
+
print(f"Sorted centroids: {centroids}")
|
|
174
|
+
return labels, {"centroids": centroids.tolist()}
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def segment_gmm(stack, num_classes):
|
|
178
|
+
"""
|
|
179
|
+
Gaussian Mixture Model segmentation.
|
|
180
|
+
Fits a GMM to the intensity histogram, then classifies each pixel
|
|
181
|
+
by posterior probability. Components sorted by mean.
|
|
182
|
+
"""
|
|
183
|
+
from sklearn.mixture import GaussianMixture
|
|
184
|
+
|
|
185
|
+
print(f"\nRunning GMM segmentation with {num_classes} components...")
|
|
186
|
+
|
|
187
|
+
# Subsample for fitting
|
|
188
|
+
flat = stack.ravel().astype(np.float32)
|
|
189
|
+
subsample_idx = np.random.choice(len(flat), size=min(2_000_000, len(flat)), replace=False)
|
|
190
|
+
subsample = flat[subsample_idx].reshape(-1, 1)
|
|
191
|
+
|
|
192
|
+
gmm = GaussianMixture(
|
|
193
|
+
n_components=num_classes, covariance_type="full", random_state=42, n_init=3, max_iter=200
|
|
194
|
+
)
|
|
195
|
+
gmm.fit(subsample)
|
|
196
|
+
|
|
197
|
+
# Sort components by mean
|
|
198
|
+
sorted_idx = np.argsort(gmm.means_.ravel())
|
|
199
|
+
label_map = np.zeros(num_classes, dtype=np.uint8)
|
|
200
|
+
for new_label, old_label in enumerate(sorted_idx):
|
|
201
|
+
label_map[old_label] = new_label
|
|
202
|
+
|
|
203
|
+
# Predict all pixels
|
|
204
|
+
print("Assigning labels to full volume...")
|
|
205
|
+
labels = np.empty(stack.shape, dtype=np.uint8)
|
|
206
|
+
for i in tqdm(range(stack.shape[0]), desc="Segmenting slices"):
|
|
207
|
+
sl = stack[i].ravel().astype(np.float32).reshape(-1, 1)
|
|
208
|
+
pred = gmm.predict(sl)
|
|
209
|
+
labels[i] = label_map[pred].reshape(stack[i].shape)
|
|
210
|
+
|
|
211
|
+
means = gmm.means_.ravel()[sorted_idx]
|
|
212
|
+
stds = np.sqrt(gmm.covariances_.ravel()[sorted_idx])
|
|
213
|
+
print(f"Sorted means: {means}")
|
|
214
|
+
print(f"Sorted stds: {stds}")
|
|
215
|
+
return labels, {
|
|
216
|
+
"means": means.tolist(),
|
|
217
|
+
"stds": stds.tolist(),
|
|
218
|
+
"weights": gmm.weights_[sorted_idx].tolist(),
|
|
219
|
+
}
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def segment_unet(stack, model_path, num_classes, device="cuda"):
|
|
223
|
+
"""
|
|
224
|
+
Supervised U-Net segmentation using a trained model.
|
|
225
|
+
"""
|
|
226
|
+
import torch
|
|
227
|
+
|
|
228
|
+
from ct_seg.model import UNetSegmentation
|
|
229
|
+
|
|
230
|
+
print(f"\nRunning U-Net segmentation with {num_classes} classes...")
|
|
231
|
+
print(f"Model: {model_path}")
|
|
232
|
+
|
|
233
|
+
dev = torch.device(device if torch.cuda.is_available() else "cpu")
|
|
234
|
+
print(f"Device: {dev}")
|
|
235
|
+
|
|
236
|
+
# Load model
|
|
237
|
+
model = UNetSegmentation(in_channels=1, num_classes=num_classes, base_filters=32)
|
|
238
|
+
# weights_only=True: restrict deserialization to tensors/containers.
|
|
239
|
+
# torch.load unpickles arbitrary Python by default, so a malicious
|
|
240
|
+
# checkpoint file would execute code on load (CWE-502 / bandit B614).
|
|
241
|
+
checkpoint = torch.load(model_path, map_location="cpu", weights_only=True)
|
|
242
|
+
if "model_state_dict" in checkpoint:
|
|
243
|
+
model.load_state_dict(checkpoint["model_state_dict"])
|
|
244
|
+
else:
|
|
245
|
+
model.load_state_dict(checkpoint)
|
|
246
|
+
model.to(dev).eval()
|
|
247
|
+
|
|
248
|
+
# Normalize
|
|
249
|
+
stack_norm = normalize_stack(stack)
|
|
250
|
+
|
|
251
|
+
# Segment slice by slice
|
|
252
|
+
labels = np.empty(stack.shape, dtype=np.uint8)
|
|
253
|
+
with torch.no_grad():
|
|
254
|
+
for i in tqdm(range(stack.shape[0]), desc="Segmenting slices"):
|
|
255
|
+
x = torch.from_numpy(stack_norm[i : i + 1][np.newaxis]).to(dev) # [1, 1, H, W]
|
|
256
|
+
pred = model.predict(x) # [1, H, W]
|
|
257
|
+
labels[i] = pred[0].cpu().numpy().astype(np.uint8)
|
|
258
|
+
|
|
259
|
+
return labels, {}
|
|
260
|
+
|
|
261
|
+
|
|
262
|
+
# =============================================================================
|
|
263
|
+
# Overlay Visualization
|
|
264
|
+
# =============================================================================
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def create_overlay(image, labels, num_classes, alpha=0.4):
|
|
268
|
+
"""Create an RGB overlay of segmentation on the original image."""
|
|
269
|
+
# Color map for classes
|
|
270
|
+
colors = [
|
|
271
|
+
[0, 0, 0], # class 0: black (background)
|
|
272
|
+
[255, 50, 50], # class 1: red
|
|
273
|
+
[50, 255, 50], # class 2: green
|
|
274
|
+
[50, 100, 255], # class 3: blue
|
|
275
|
+
[255, 255, 50], # class 4: yellow
|
|
276
|
+
[255, 50, 255], # class 5: magenta
|
|
277
|
+
[50, 255, 255], # class 6: cyan
|
|
278
|
+
[255, 165, 0], # class 7: orange
|
|
279
|
+
]
|
|
280
|
+
|
|
281
|
+
# Normalize image to [0, 255]
|
|
282
|
+
img = image.astype(np.float32)
|
|
283
|
+
img = (img - img.min()) / (img.max() - img.min() + 1e-8) * 255
|
|
284
|
+
img_rgb = np.stack([img, img, img], axis=-1).astype(np.uint8)
|
|
285
|
+
|
|
286
|
+
# Create colored mask
|
|
287
|
+
mask_rgb = np.zeros((*labels.shape, 3), dtype=np.uint8)
|
|
288
|
+
for c in range(min(num_classes, len(colors))):
|
|
289
|
+
mask_rgb[labels == c] = colors[c]
|
|
290
|
+
|
|
291
|
+
# Blend
|
|
292
|
+
overlay = (img_rgb * (1 - alpha) + mask_rgb * alpha).astype(np.uint8)
|
|
293
|
+
return overlay
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
def save_overlays(stack, labels, out_dir, num_classes, every_n=10, offset=0):
|
|
297
|
+
"""Save overlay images for every Nth slice."""
|
|
298
|
+
from PIL import Image
|
|
299
|
+
|
|
300
|
+
out_dir = Path(out_dir)
|
|
301
|
+
out_dir.mkdir(parents=True, exist_ok=True)
|
|
302
|
+
|
|
303
|
+
for i in range(0, stack.shape[0], every_n):
|
|
304
|
+
overlay = create_overlay(stack[i], labels[i], num_classes)
|
|
305
|
+
Image.fromarray(overlay).save(str(out_dir / f"overlay_{i + offset:05d}.png"))
|
|
306
|
+
print(f"Saved {stack.shape[0] // every_n} overlay images to {out_dir}")
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
# =============================================================================
|
|
310
|
+
# Statistics
|
|
311
|
+
# =============================================================================
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def compute_statistics(labels, num_classes):
|
|
315
|
+
"""Compute volume fraction and voxel count per class."""
|
|
316
|
+
total = labels.size
|
|
317
|
+
print(f"\n{'='*50}")
|
|
318
|
+
print("Segmentation Statistics")
|
|
319
|
+
print(f"{'='*50}")
|
|
320
|
+
print(f"{'Class':<10} {'Voxels':>15} {'Volume %':>12}")
|
|
321
|
+
print(f"{'-'*37}")
|
|
322
|
+
stats = {}
|
|
323
|
+
for c in range(num_classes):
|
|
324
|
+
count = int(np.sum(labels == c))
|
|
325
|
+
frac = count / total * 100
|
|
326
|
+
print(f"{c:<10} {count:>15,} {frac:>11.2f}%")
|
|
327
|
+
stats[f"class_{c}"] = {"voxels": count, "volume_fraction": frac}
|
|
328
|
+
print(f"{'-'*37}")
|
|
329
|
+
print(f"{'Total':<10} {total:>15,} {'100.00%':>12}")
|
|
330
|
+
return stats
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
# =============================================================================
|
|
334
|
+
# Main
|
|
335
|
+
# =============================================================================
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def main():
|
|
339
|
+
parser = argparse.ArgumentParser(description="CT Image Stack Segmentation")
|
|
340
|
+
parser.add_argument("--input", required=True, help="Path to directory of tiff files")
|
|
341
|
+
parser.add_argument(
|
|
342
|
+
"--output", default=None, help="Output directory (default: input_segmented)"
|
|
343
|
+
)
|
|
344
|
+
parser.add_argument(
|
|
345
|
+
"--method",
|
|
346
|
+
required=True,
|
|
347
|
+
choices=["otsu", "kmeans", "gmm", "unet"],
|
|
348
|
+
help="Segmentation method",
|
|
349
|
+
)
|
|
350
|
+
parser.add_argument(
|
|
351
|
+
"--num_classes", type=int, default=4, help="Number of classes/phases (default: 4)"
|
|
352
|
+
)
|
|
353
|
+
parser.add_argument(
|
|
354
|
+
"--model", type=str, default=None, help="Path to trained U-Net model (for unet method)"
|
|
355
|
+
)
|
|
356
|
+
parser.add_argument(
|
|
357
|
+
"--slice_range",
|
|
358
|
+
nargs=2,
|
|
359
|
+
type=int,
|
|
360
|
+
default=None,
|
|
361
|
+
metavar=("START", "END"),
|
|
362
|
+
help="Process only a range of slices",
|
|
363
|
+
)
|
|
364
|
+
parser.add_argument(
|
|
365
|
+
"--enhance", action="store_true", help="Apply contrast enhancement before segmentation"
|
|
366
|
+
)
|
|
367
|
+
parser.add_argument("--save_overlay", action="store_true", help="Save overlay images")
|
|
368
|
+
parser.add_argument(
|
|
369
|
+
"--overlay_every", type=int, default=10, help="Save overlay every N slices (default: 10)"
|
|
370
|
+
)
|
|
371
|
+
parser.add_argument(
|
|
372
|
+
"--device", type=str, default="cuda", help="Device for U-Net (default: cuda)"
|
|
373
|
+
)
|
|
374
|
+
args = parser.parse_args()
|
|
375
|
+
|
|
376
|
+
# Validate
|
|
377
|
+
if args.method == "unet" and args.model is None:
|
|
378
|
+
parser.error("--model is required for unet method")
|
|
379
|
+
|
|
380
|
+
# Setup output directory
|
|
381
|
+
if args.output is None:
|
|
382
|
+
args.output = str(
|
|
383
|
+
Path(args.input).parent / (Path(args.input).name + f"_segmented_{args.method}")
|
|
384
|
+
)
|
|
385
|
+
|
|
386
|
+
# Load data
|
|
387
|
+
print(f"\n{'='*60}")
|
|
388
|
+
print("CT Image Stack Segmentation")
|
|
389
|
+
print(f"{'='*60}")
|
|
390
|
+
print(f"Input: {args.input}")
|
|
391
|
+
print(f"Method: {args.method}")
|
|
392
|
+
print(f"Classes: {args.num_classes}")
|
|
393
|
+
print(f"Output: {args.output}")
|
|
394
|
+
|
|
395
|
+
start = args.slice_range[0] if args.slice_range else None
|
|
396
|
+
end = args.slice_range[1] if args.slice_range else None
|
|
397
|
+
stack, paths = load_image_stack(args.input, start=start, end=end)
|
|
398
|
+
print(f"Loaded: {stack.shape} ({stack.dtype})")
|
|
399
|
+
|
|
400
|
+
# Preprocess
|
|
401
|
+
if args.enhance:
|
|
402
|
+
print("Applying contrast enhancement...")
|
|
403
|
+
stack_proc = enhance_contrast(stack)
|
|
404
|
+
else:
|
|
405
|
+
stack_proc = normalize_stack(stack)
|
|
406
|
+
|
|
407
|
+
# Segment
|
|
408
|
+
if args.method == "otsu":
|
|
409
|
+
labels, info = segment_otsu(stack_proc, args.num_classes)
|
|
410
|
+
elif args.method == "kmeans":
|
|
411
|
+
labels, info = segment_kmeans(stack_proc, args.num_classes)
|
|
412
|
+
elif args.method == "gmm":
|
|
413
|
+
labels, info = segment_gmm(stack_proc, args.num_classes)
|
|
414
|
+
elif args.method == "unet":
|
|
415
|
+
labels, info = segment_unet(stack, args.model, args.num_classes, args.device)
|
|
416
|
+
|
|
417
|
+
# Statistics
|
|
418
|
+
stats = compute_statistics(labels, args.num_classes)
|
|
419
|
+
|
|
420
|
+
# Save segmented volume (colored PNGs + raw label tiffs for 3D viewer)
|
|
421
|
+
offset = start if start else 0
|
|
422
|
+
save_tiff_stack(labels, args.output, args.num_classes, offset=offset)
|
|
423
|
+
|
|
424
|
+
# Also save raw labels as tiffs for viewer/3D rendering
|
|
425
|
+
raw_label_dir = Path(args.output) / "raw_labels"
|
|
426
|
+
raw_label_dir.mkdir(parents=True, exist_ok=True)
|
|
427
|
+
for i in range(labels.shape[0]):
|
|
428
|
+
tifffile.imwrite(str(raw_label_dir / f"{i + offset:05d}.tiff"), labels[i])
|
|
429
|
+
|
|
430
|
+
# Save overlays
|
|
431
|
+
if args.save_overlay:
|
|
432
|
+
overlay_dir = str(Path(args.output) / "overlays")
|
|
433
|
+
save_overlays(
|
|
434
|
+
stack, labels, overlay_dir, args.num_classes, every_n=args.overlay_every, offset=offset
|
|
435
|
+
)
|
|
436
|
+
|
|
437
|
+
# Save metadata
|
|
438
|
+
import json
|
|
439
|
+
|
|
440
|
+
meta = {
|
|
441
|
+
"method": args.method,
|
|
442
|
+
"num_classes": args.num_classes,
|
|
443
|
+
"input": args.input,
|
|
444
|
+
"shape": list(stack.shape),
|
|
445
|
+
"dtype": str(stack.dtype),
|
|
446
|
+
"enhanced": args.enhance,
|
|
447
|
+
"info": info,
|
|
448
|
+
"statistics": stats,
|
|
449
|
+
}
|
|
450
|
+
meta_path = Path(args.output) / "segmentation_info.json"
|
|
451
|
+
with open(meta_path, "w") as f:
|
|
452
|
+
json.dump(meta, f, indent=2)
|
|
453
|
+
print(f"\nMetadata saved to {meta_path}")
|
|
454
|
+
|
|
455
|
+
print(f"\n{'='*60}")
|
|
456
|
+
print("Segmentation complete!")
|
|
457
|
+
print(f"{'='*60}")
|
|
458
|
+
|
|
459
|
+
|
|
460
|
+
if __name__ == "__main__":
|
|
461
|
+
main()
|