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/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()