ct-segmentation-toolkit 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,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)
@@ -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
@@ -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