torchtomo 0.1.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.
- torchtomo/__init__.py +35 -0
- torchtomo/base.py +89 -0
- torchtomo/fanbeam.py +383 -0
- torchtomo/filters.py +102 -0
- torchtomo/parallel.py +217 -0
- torchtomo/phantom.py +186 -0
- torchtomo-0.1.0.dist-info/METADATA +152 -0
- torchtomo-0.1.0.dist-info/RECORD +11 -0
- torchtomo-0.1.0.dist-info/WHEEL +5 -0
- torchtomo-0.1.0.dist-info/licenses/LICENSE +12 -0
- torchtomo-0.1.0.dist-info/top_level.txt +1 -0
torchtomo/__init__.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""
|
|
2
|
+
TorchTomo: Differentiable CT Reconstruction in Pure PyTorch
|
|
3
|
+
|
|
4
|
+
A lightweight library for CT forward and back projection that works
|
|
5
|
+
on any device (CPU, CUDA, MPS) without compilation.
|
|
6
|
+
|
|
7
|
+
Example:
|
|
8
|
+
>>> from torchtomo import ParallelBeam, FanBeam
|
|
9
|
+
>>>
|
|
10
|
+
>>> # Parallel beam
|
|
11
|
+
>>> projector = ParallelBeam(img_size=256, n_angles=180, n_det=256)
|
|
12
|
+
>>> sinogram = projector.forward(image)
|
|
13
|
+
>>> recon = projector.fbp(sinogram)
|
|
14
|
+
>>>
|
|
15
|
+
>>> # Fan beam
|
|
16
|
+
>>> projector = FanBeam(img_size=256, n_angles=360, n_det=400,
|
|
17
|
+
... src_dist=500, det_dist=500)
|
|
18
|
+
>>> sinogram = projector.forward(image)
|
|
19
|
+
>>> recon = projector.fbp(sinogram)
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from .fanbeam import FanBeam
|
|
23
|
+
from .filters import apply_filter, get_filter
|
|
24
|
+
from .parallel import ParallelBeam
|
|
25
|
+
from .phantom import circle_phantom, shepp_logan
|
|
26
|
+
|
|
27
|
+
__version__ = "0.1.0"
|
|
28
|
+
__all__ = [
|
|
29
|
+
"ParallelBeam",
|
|
30
|
+
"FanBeam",
|
|
31
|
+
"apply_filter",
|
|
32
|
+
"get_filter",
|
|
33
|
+
"shepp_logan",
|
|
34
|
+
"circle_phantom",
|
|
35
|
+
]
|
torchtomo/base.py
ADDED
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
"""Base class for CT projectors."""
|
|
2
|
+
|
|
3
|
+
from abc import ABC, abstractmethod
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
import torch.nn as nn
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class BaseProjector(nn.Module, ABC):
|
|
10
|
+
"""
|
|
11
|
+
Abstract base class for CT projectors.
|
|
12
|
+
|
|
13
|
+
All projectors support:
|
|
14
|
+
- forward(): Image -> Sinogram (Radon transform)
|
|
15
|
+
- backward(): Sinogram -> Image (Adjoint/back-projection)
|
|
16
|
+
- fbp(): Filtered back-projection reconstruction
|
|
17
|
+
|
|
18
|
+
All operations are differentiable.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
def __init__(
|
|
22
|
+
self,
|
|
23
|
+
img_size: int,
|
|
24
|
+
n_angles: int,
|
|
25
|
+
n_det: int,
|
|
26
|
+
angle_range: tuple[float, float] = (0, torch.pi),
|
|
27
|
+
):
|
|
28
|
+
super().__init__()
|
|
29
|
+
self.img_size = img_size
|
|
30
|
+
self.n_angles = n_angles
|
|
31
|
+
self.n_det = n_det
|
|
32
|
+
self.angle_range = angle_range
|
|
33
|
+
|
|
34
|
+
angles = torch.linspace(
|
|
35
|
+
angle_range[0], angle_range[1], n_angles, dtype=torch.float32
|
|
36
|
+
)
|
|
37
|
+
self.register_buffer("angles", angles)
|
|
38
|
+
|
|
39
|
+
@abstractmethod
|
|
40
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
41
|
+
"""
|
|
42
|
+
Forward projection: image -> sinogram.
|
|
43
|
+
|
|
44
|
+
Args:
|
|
45
|
+
x: Image tensor of shape [B, 1, H, W]
|
|
46
|
+
|
|
47
|
+
Returns:
|
|
48
|
+
Sinogram of shape [B, 1, n_angles, n_det]
|
|
49
|
+
"""
|
|
50
|
+
pass
|
|
51
|
+
|
|
52
|
+
@abstractmethod
|
|
53
|
+
def backward(self, sinogram: torch.Tensor) -> torch.Tensor:
|
|
54
|
+
"""
|
|
55
|
+
Back projection (adjoint): sinogram -> image.
|
|
56
|
+
|
|
57
|
+
Note: This is NOT the inverse, just the adjoint operator.
|
|
58
|
+
For reconstruction, use fbp().
|
|
59
|
+
|
|
60
|
+
Args:
|
|
61
|
+
sinogram: Sinogram of shape [B, 1, n_angles, n_det]
|
|
62
|
+
|
|
63
|
+
Returns:
|
|
64
|
+
Back-projected image of shape [B, 1, H, W]
|
|
65
|
+
"""
|
|
66
|
+
pass
|
|
67
|
+
|
|
68
|
+
@abstractmethod
|
|
69
|
+
def fbp(self, sinogram: torch.Tensor, filter_name: str = "ramp") -> torch.Tensor:
|
|
70
|
+
"""
|
|
71
|
+
Filtered back-projection reconstruction.
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
sinogram: Sinogram of shape [B, 1, n_angles, n_det]
|
|
75
|
+
filter_name: Filter type ('ramp', 'shepp-logan', 'cosine',
|
|
76
|
+
'hamming', 'hann')
|
|
77
|
+
|
|
78
|
+
Returns:
|
|
79
|
+
Reconstructed image of shape [B, 1, H, W]
|
|
80
|
+
"""
|
|
81
|
+
pass
|
|
82
|
+
|
|
83
|
+
def __repr__(self) -> str:
|
|
84
|
+
return (
|
|
85
|
+
f"{self.__class__.__name__}("
|
|
86
|
+
f"img_size={self.img_size}, "
|
|
87
|
+
f"n_angles={self.n_angles}, "
|
|
88
|
+
f"n_det={self.n_det})"
|
|
89
|
+
)
|
torchtomo/fanbeam.py
ADDED
|
@@ -0,0 +1,383 @@
|
|
|
1
|
+
"""Fan beam CT projector with flat detector."""
|
|
2
|
+
|
|
3
|
+
from typing import Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
import torch.nn.functional as F
|
|
8
|
+
|
|
9
|
+
from .base import BaseProjector
|
|
10
|
+
from .filters import FilterType, apply_filter
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class FanBeam(BaseProjector):
|
|
14
|
+
"""
|
|
15
|
+
Fan beam CT projector with flat detector.
|
|
16
|
+
|
|
17
|
+
In fan beam geometry, X-rays emanate from a point source and spread
|
|
18
|
+
out in a fan shape to a flat detector array. This is the standard geometry
|
|
19
|
+
for clinical CT scanners.
|
|
20
|
+
|
|
21
|
+
Geometry:
|
|
22
|
+
- Point source at distance src_dist from origin
|
|
23
|
+
- Flat detector at distance det_dist from origin (opposite side)
|
|
24
|
+
- Source and detector rotate around the object
|
|
25
|
+
|
|
26
|
+
Example:
|
|
27
|
+
>>> projector = FanBeam(
|
|
28
|
+
... img_size=256,
|
|
29
|
+
... n_angles=360,
|
|
30
|
+
... n_det=400,
|
|
31
|
+
... src_dist=500,
|
|
32
|
+
... det_dist=500
|
|
33
|
+
... )
|
|
34
|
+
>>> sinogram = projector.forward(image)
|
|
35
|
+
>>> recon = projector.fbp(sinogram)
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
def __init__(
|
|
39
|
+
self,
|
|
40
|
+
img_size: int = 256,
|
|
41
|
+
n_angles: int = 360,
|
|
42
|
+
n_det: int = 400,
|
|
43
|
+
src_dist: float = 500.0,
|
|
44
|
+
det_dist: float = 500.0,
|
|
45
|
+
det_width: Optional[float] = None,
|
|
46
|
+
det_spacing: Optional[float] = None,
|
|
47
|
+
angle_range: tuple[float, float] = (0, 2 * np.pi),
|
|
48
|
+
n_samples: int = 512,
|
|
49
|
+
circle: bool = True,
|
|
50
|
+
):
|
|
51
|
+
"""
|
|
52
|
+
Initialize fan beam projector.
|
|
53
|
+
|
|
54
|
+
Args:
|
|
55
|
+
img_size: Image size (assumed square)
|
|
56
|
+
n_angles: Number of projection angles
|
|
57
|
+
n_det: Number of detector elements
|
|
58
|
+
src_dist: Source to isocenter distance (in pixels)
|
|
59
|
+
det_dist: Isocenter to detector distance (in pixels)
|
|
60
|
+
det_width: Total detector width (alternative to det_spacing)
|
|
61
|
+
det_spacing: Spacing between detector elements (alternative to det_width)
|
|
62
|
+
angle_range: Range of angles (default: full rotation)
|
|
63
|
+
n_samples: Number of samples per ray for integration
|
|
64
|
+
circle: If True, mask image to inscribed circle
|
|
65
|
+
"""
|
|
66
|
+
super().__init__(img_size, n_angles, n_det, angle_range)
|
|
67
|
+
|
|
68
|
+
self.src_dist = src_dist
|
|
69
|
+
self.det_dist = det_dist
|
|
70
|
+
magnification = (src_dist + det_dist) / src_dist
|
|
71
|
+
if det_spacing is not None:
|
|
72
|
+
self.det_width = det_spacing * n_det
|
|
73
|
+
elif det_width is not None:
|
|
74
|
+
self.det_width = det_width
|
|
75
|
+
else:
|
|
76
|
+
self.det_width = 1.5 * magnification * img_size
|
|
77
|
+
self.n_samples = n_samples
|
|
78
|
+
self.circle = circle
|
|
79
|
+
|
|
80
|
+
self.scale = 2.0 / img_size
|
|
81
|
+
self._src_dist_norm = src_dist * self.scale
|
|
82
|
+
self._det_dist_norm = det_dist * self.scale
|
|
83
|
+
self._det_width_norm = self.det_width * self.scale
|
|
84
|
+
|
|
85
|
+
ray_grids, ray_lengths = self._precompute_ray_grids()
|
|
86
|
+
self.register_buffer("ray_grids", ray_grids)
|
|
87
|
+
self.register_buffer("ray_lengths", ray_lengths)
|
|
88
|
+
|
|
89
|
+
back_grids, weights = self._precompute_backward_grids()
|
|
90
|
+
self.register_buffer("backward_grids", back_grids)
|
|
91
|
+
self.register_buffer("backward_weights", weights)
|
|
92
|
+
|
|
93
|
+
det_pos = torch.linspace(
|
|
94
|
+
-self._det_width_norm / 2, self._det_width_norm / 2, n_det
|
|
95
|
+
)
|
|
96
|
+
D = self._src_dist_norm + self._det_dist_norm
|
|
97
|
+
cos_weight = D / torch.sqrt(D**2 + det_pos**2)
|
|
98
|
+
self.register_buffer("cos_weight", cos_weight)
|
|
99
|
+
|
|
100
|
+
if circle:
|
|
101
|
+
coords = torch.linspace(-1, 1, img_size)
|
|
102
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
103
|
+
mask = (x**2 + y**2 <= 1).float()
|
|
104
|
+
self.register_buffer("circle_mask", mask)
|
|
105
|
+
|
|
106
|
+
def _precompute_ray_grids(self) -> tuple[torch.Tensor, torch.Tensor]:
|
|
107
|
+
"""
|
|
108
|
+
Precompute sampling grids for all rays.
|
|
109
|
+
|
|
110
|
+
Returns:
|
|
111
|
+
grids: [n_angles, n_det, n_samples, 2]
|
|
112
|
+
ray_lengths: [n_angles, n_det] path length through image for each ray
|
|
113
|
+
"""
|
|
114
|
+
all_grids = []
|
|
115
|
+
all_lengths = []
|
|
116
|
+
|
|
117
|
+
for angle in self.angles:
|
|
118
|
+
grid, lengths = self._compute_rays_for_angle(angle)
|
|
119
|
+
all_grids.append(grid)
|
|
120
|
+
all_lengths.append(lengths)
|
|
121
|
+
|
|
122
|
+
return torch.stack(all_grids), torch.stack(all_lengths)
|
|
123
|
+
|
|
124
|
+
def _compute_rays_for_angle(
|
|
125
|
+
self, angle: torch.Tensor
|
|
126
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
127
|
+
"""
|
|
128
|
+
Compute ray sampling points for a single angle.
|
|
129
|
+
|
|
130
|
+
Only samples the portion of each ray that intersects the unit circle
|
|
131
|
+
(image region in normalized coordinates).
|
|
132
|
+
|
|
133
|
+
Returns:
|
|
134
|
+
grid: [n_det, n_samples, 2]
|
|
135
|
+
ray_lengths: [n_det] path length through image for each ray
|
|
136
|
+
"""
|
|
137
|
+
cos_a = torch.cos(angle)
|
|
138
|
+
sin_a = torch.sin(angle)
|
|
139
|
+
|
|
140
|
+
src_x = -self._src_dist_norm * sin_a
|
|
141
|
+
src_y = self._src_dist_norm * cos_a
|
|
142
|
+
|
|
143
|
+
det_cx = self._det_dist_norm * sin_a
|
|
144
|
+
det_cy = -self._det_dist_norm * cos_a
|
|
145
|
+
|
|
146
|
+
det_dir_x = cos_a
|
|
147
|
+
det_dir_y = sin_a
|
|
148
|
+
|
|
149
|
+
det_offsets = torch.linspace(
|
|
150
|
+
-self._det_width_norm / 2, self._det_width_norm / 2, self.n_det
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
det_x = det_cx + det_offsets * det_dir_x
|
|
154
|
+
det_y = det_cy + det_offsets * det_dir_y
|
|
155
|
+
|
|
156
|
+
dir_x = det_x - src_x
|
|
157
|
+
dir_y = det_y - src_y
|
|
158
|
+
ray_len_full = torch.sqrt(dir_x**2 + dir_y**2)
|
|
159
|
+
dir_x = dir_x / ray_len_full
|
|
160
|
+
dir_y = dir_y / ray_len_full
|
|
161
|
+
|
|
162
|
+
a = dir_x**2 + dir_y**2
|
|
163
|
+
b = 2 * (src_x * dir_x + src_y * dir_y)
|
|
164
|
+
c = src_x**2 + src_y**2 - 1.0
|
|
165
|
+
|
|
166
|
+
discriminant = b**2 - 4 * a * c
|
|
167
|
+
discriminant = torch.clamp(discriminant, min=0)
|
|
168
|
+
|
|
169
|
+
sqrt_disc = torch.sqrt(discriminant)
|
|
170
|
+
t_entry = (-b - sqrt_disc) / (2 * a)
|
|
171
|
+
t_exit = (-b + sqrt_disc) / (2 * a)
|
|
172
|
+
|
|
173
|
+
t_entry = torch.clamp(t_entry, min=0)
|
|
174
|
+
t_exit = torch.clamp(t_exit, min=t_entry)
|
|
175
|
+
|
|
176
|
+
ray_lengths = t_exit - t_entry
|
|
177
|
+
|
|
178
|
+
t_samples = torch.linspace(0, 1, self.n_samples).view(1, -1)
|
|
179
|
+
t_actual = t_entry.view(-1, 1) + t_samples * (t_exit - t_entry).view(-1, 1)
|
|
180
|
+
|
|
181
|
+
ray_x = src_x + t_actual * dir_x.view(-1, 1)
|
|
182
|
+
ray_y = src_y + t_actual * dir_y.view(-1, 1)
|
|
183
|
+
|
|
184
|
+
grid = torch.stack([ray_x, ray_y], dim=-1)
|
|
185
|
+
|
|
186
|
+
return grid, ray_lengths
|
|
187
|
+
|
|
188
|
+
def _precompute_backward_grids(self) -> tuple[torch.Tensor, torch.Tensor]:
|
|
189
|
+
"""
|
|
190
|
+
Precompute grids and weights for back-projection.
|
|
191
|
+
|
|
192
|
+
Returns:
|
|
193
|
+
grids: [n_angles, H, W, 2] sampling positions in sinogram
|
|
194
|
+
weights: [n_angles, H, W] distance weighting
|
|
195
|
+
"""
|
|
196
|
+
grids = []
|
|
197
|
+
weights = []
|
|
198
|
+
|
|
199
|
+
coords = torch.linspace(-1, 1, self.img_size)
|
|
200
|
+
grid_y, grid_x = torch.meshgrid(coords, coords, indexing="ij")
|
|
201
|
+
|
|
202
|
+
for angle in self.angles:
|
|
203
|
+
grid, weight = self._compute_backward_for_angle(angle, grid_x, grid_y)
|
|
204
|
+
grids.append(grid)
|
|
205
|
+
weights.append(weight)
|
|
206
|
+
|
|
207
|
+
return torch.stack(grids), torch.stack(weights)
|
|
208
|
+
|
|
209
|
+
def _compute_backward_for_angle(
|
|
210
|
+
self,
|
|
211
|
+
angle: torch.Tensor,
|
|
212
|
+
grid_x: torch.Tensor,
|
|
213
|
+
grid_y: torch.Tensor,
|
|
214
|
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
215
|
+
"""
|
|
216
|
+
Compute back-projection mapping for a single angle.
|
|
217
|
+
|
|
218
|
+
For each pixel, determine which detector element it maps to and
|
|
219
|
+
compute the FBP weight U² where U = D / (D + x*sin(β) - y*cos(β)).
|
|
220
|
+
"""
|
|
221
|
+
cos_a = torch.cos(angle)
|
|
222
|
+
sin_a = torch.sin(angle)
|
|
223
|
+
|
|
224
|
+
src_x = -self._src_dist_norm * sin_a
|
|
225
|
+
src_y = self._src_dist_norm * cos_a
|
|
226
|
+
|
|
227
|
+
det_cx = self._det_dist_norm * sin_a
|
|
228
|
+
det_cy = -self._det_dist_norm * cos_a
|
|
229
|
+
|
|
230
|
+
px_x = grid_x - src_x
|
|
231
|
+
px_y = grid_y - src_y
|
|
232
|
+
|
|
233
|
+
sd_x = det_cx - src_x
|
|
234
|
+
sd_y = det_cy - src_y
|
|
235
|
+
sd_len = torch.sqrt(sd_x**2 + sd_y**2)
|
|
236
|
+
|
|
237
|
+
sd_ux = sd_x / sd_len
|
|
238
|
+
sd_uy = sd_y / sd_len
|
|
239
|
+
|
|
240
|
+
proj_len = px_x * sd_ux + px_y * sd_uy
|
|
241
|
+
|
|
242
|
+
t = sd_len / (proj_len + 1e-8)
|
|
243
|
+
|
|
244
|
+
int_x = src_x + t * px_x
|
|
245
|
+
int_y = src_y + t * px_y
|
|
246
|
+
|
|
247
|
+
det_dir_x = cos_a
|
|
248
|
+
det_dir_y = sin_a
|
|
249
|
+
|
|
250
|
+
det_offset = (int_x - det_cx) * det_dir_x + (int_y - det_cy) * det_dir_y
|
|
251
|
+
|
|
252
|
+
det_normalized = det_offset / (self._det_width_norm / 2)
|
|
253
|
+
|
|
254
|
+
grid = torch.zeros(self.img_size, self.img_size, 2)
|
|
255
|
+
grid[..., 0] = det_normalized
|
|
256
|
+
grid[..., 1] = 0
|
|
257
|
+
|
|
258
|
+
D = self._src_dist_norm + self._det_dist_norm
|
|
259
|
+
U = D / (D + grid_x * sin_a + grid_y * cos_a)
|
|
260
|
+
weight = U * U
|
|
261
|
+
|
|
262
|
+
return grid, weight
|
|
263
|
+
|
|
264
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
265
|
+
"""
|
|
266
|
+
Forward projection: image -> sinogram.
|
|
267
|
+
|
|
268
|
+
Args:
|
|
269
|
+
x: Image tensor [B, 1, H, W]
|
|
270
|
+
|
|
271
|
+
Returns:
|
|
272
|
+
Sinogram [B, 1, n_angles, n_det]
|
|
273
|
+
"""
|
|
274
|
+
B = x.shape[0]
|
|
275
|
+
device = x.device
|
|
276
|
+
|
|
277
|
+
if self.circle:
|
|
278
|
+
x = x * self.circle_mask.view(1, 1, self.img_size, self.img_size)
|
|
279
|
+
|
|
280
|
+
projections = []
|
|
281
|
+
|
|
282
|
+
for i in range(self.n_angles):
|
|
283
|
+
grid = self.ray_grids[i]
|
|
284
|
+
ray_len = self.ray_lengths[i]
|
|
285
|
+
|
|
286
|
+
grid = grid.unsqueeze(0).expand(B, -1, -1, -1).to(device)
|
|
287
|
+
|
|
288
|
+
samples = F.grid_sample(
|
|
289
|
+
x,
|
|
290
|
+
grid,
|
|
291
|
+
mode="bilinear",
|
|
292
|
+
padding_mode="zeros",
|
|
293
|
+
align_corners=True,
|
|
294
|
+
)
|
|
295
|
+
|
|
296
|
+
projection = samples.mean(dim=-1)
|
|
297
|
+
projection = projection * ray_len.view(1, 1, -1).to(device)
|
|
298
|
+
projections.append(projection)
|
|
299
|
+
|
|
300
|
+
sinogram = torch.stack(projections, dim=2)
|
|
301
|
+
|
|
302
|
+
return sinogram
|
|
303
|
+
|
|
304
|
+
def backward(self, sinogram: torch.Tensor) -> torch.Tensor:
|
|
305
|
+
"""
|
|
306
|
+
Back projection (adjoint): sinogram -> image.
|
|
307
|
+
|
|
308
|
+
Args:
|
|
309
|
+
sinogram: Sinogram [B, 1, n_angles, n_det]
|
|
310
|
+
|
|
311
|
+
Returns:
|
|
312
|
+
Back-projected image [B, 1, H, W]
|
|
313
|
+
"""
|
|
314
|
+
B = sinogram.shape[0]
|
|
315
|
+
device = sinogram.device
|
|
316
|
+
|
|
317
|
+
recon = torch.zeros(B, 1, self.img_size, self.img_size, device=device)
|
|
318
|
+
|
|
319
|
+
for i in range(self.n_angles):
|
|
320
|
+
# Sinogram row [B, 1, 1, n_det]
|
|
321
|
+
sino_row = sinogram[:, :, i : i + 1, :]
|
|
322
|
+
|
|
323
|
+
# Sampling grid
|
|
324
|
+
grid = self.backward_grids[i].unsqueeze(0).expand(B, -1, -1, -1)
|
|
325
|
+
grid = grid.to(device)
|
|
326
|
+
|
|
327
|
+
# Sample
|
|
328
|
+
contribution = F.grid_sample(
|
|
329
|
+
sino_row,
|
|
330
|
+
grid,
|
|
331
|
+
mode="bilinear",
|
|
332
|
+
padding_mode="zeros",
|
|
333
|
+
align_corners=True,
|
|
334
|
+
)
|
|
335
|
+
|
|
336
|
+
# Apply distance weighting
|
|
337
|
+
weight = self.backward_weights[i].view(1, 1, self.img_size, self.img_size)
|
|
338
|
+
weight = weight.to(device)
|
|
339
|
+
|
|
340
|
+
recon += contribution * weight
|
|
341
|
+
|
|
342
|
+
delta_beta = (self.angle_range[1] - self.angle_range[0]) / self.n_angles
|
|
343
|
+
recon = recon * delta_beta / 2
|
|
344
|
+
|
|
345
|
+
# Circle mask
|
|
346
|
+
if self.circle:
|
|
347
|
+
recon = recon * self.circle_mask.view(1, 1, self.img_size, self.img_size)
|
|
348
|
+
|
|
349
|
+
return recon
|
|
350
|
+
|
|
351
|
+
def fbp(
|
|
352
|
+
self, sinogram: torch.Tensor, filter_name: FilterType = "ramp"
|
|
353
|
+
) -> torch.Tensor:
|
|
354
|
+
"""
|
|
355
|
+
Filtered back-projection for fan beam with flat detector.
|
|
356
|
+
|
|
357
|
+
Args:
|
|
358
|
+
sinogram: Sinogram [B, 1, n_angles, n_det]
|
|
359
|
+
filter_name: Filter type
|
|
360
|
+
|
|
361
|
+
Returns:
|
|
362
|
+
Reconstructed image [B, 1, H, W]
|
|
363
|
+
"""
|
|
364
|
+
cos_w = self.cos_weight.view(1, 1, 1, -1).to(sinogram.device)
|
|
365
|
+
weighted_sino = sinogram * cos_w
|
|
366
|
+
|
|
367
|
+
filtered_sino = apply_filter(weighted_sino, filter_name)
|
|
368
|
+
|
|
369
|
+
filtered_sino = filtered_sino * (self.img_size / 2)
|
|
370
|
+
|
|
371
|
+
recon = self.backward(filtered_sino)
|
|
372
|
+
|
|
373
|
+
return recon
|
|
374
|
+
|
|
375
|
+
def __repr__(self) -> str:
|
|
376
|
+
return (
|
|
377
|
+
f"FanBeam("
|
|
378
|
+
f"img_size={self.img_size}, "
|
|
379
|
+
f"n_angles={self.n_angles}, "
|
|
380
|
+
f"n_det={self.n_det}, "
|
|
381
|
+
f"src_dist={self.src_dist}, "
|
|
382
|
+
f"det_dist={self.det_dist})"
|
|
383
|
+
)
|
torchtomo/filters.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
"""FBP filters for CT reconstruction."""
|
|
2
|
+
|
|
3
|
+
from typing import Literal
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
import torch.fft as fft
|
|
8
|
+
|
|
9
|
+
FilterType = Literal["ramp", "shepp-logan", "cosine", "hamming", "hann", "none"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def get_filter(
|
|
13
|
+
size: int,
|
|
14
|
+
filter_name: FilterType = "ramp",
|
|
15
|
+
device: torch.device = None,
|
|
16
|
+
dtype: torch.dtype = torch.float32,
|
|
17
|
+
) -> torch.Tensor:
|
|
18
|
+
"""
|
|
19
|
+
Generate frequency-domain filter for FBP.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
size: Filter size
|
|
23
|
+
filter_name: Type of filter
|
|
24
|
+
device: Target device
|
|
25
|
+
dtype: Data type
|
|
26
|
+
|
|
27
|
+
Returns:
|
|
28
|
+
Filter in frequency domain, shape [size]
|
|
29
|
+
"""
|
|
30
|
+
if filter_name == "none":
|
|
31
|
+
return torch.ones(size, device=device, dtype=dtype)
|
|
32
|
+
|
|
33
|
+
# Frequency axis: fftfreq gives [-0.5, 0.5) normalized frequencies
|
|
34
|
+
freq = np.fft.fftfreq(size).astype(np.float64)
|
|
35
|
+
|
|
36
|
+
# Ramp filter: |f| scaled to detector spacing
|
|
37
|
+
# In FBP, the filter is |omega| = 2*pi*|f|, but with discrete sampling
|
|
38
|
+
# we use |f| directly and scale appropriately
|
|
39
|
+
ramp = np.abs(freq)
|
|
40
|
+
|
|
41
|
+
if filter_name == "ramp":
|
|
42
|
+
filt = ramp
|
|
43
|
+
elif filter_name == "shepp-logan":
|
|
44
|
+
# sinc window (avoids division by zero)
|
|
45
|
+
with np.errstate(divide="ignore", invalid="ignore"):
|
|
46
|
+
window = np.sinc(2 * freq) # sinc(2f) = sin(2*pi*f)/(2*pi*f)
|
|
47
|
+
filt = ramp * window
|
|
48
|
+
elif filter_name == "cosine":
|
|
49
|
+
filt = ramp * np.cos(np.pi * freq)
|
|
50
|
+
elif filter_name == "hamming":
|
|
51
|
+
filt = ramp * (0.54 + 0.46 * np.cos(2 * np.pi * freq))
|
|
52
|
+
elif filter_name == "hann":
|
|
53
|
+
filt = ramp * (0.5 + 0.5 * np.cos(2 * np.pi * freq))
|
|
54
|
+
else:
|
|
55
|
+
raise ValueError(f"Unknown filter: {filter_name}")
|
|
56
|
+
|
|
57
|
+
# Convert to torch tensor
|
|
58
|
+
filt = torch.from_numpy(filt.astype(np.float32)).to(device=device, dtype=dtype)
|
|
59
|
+
|
|
60
|
+
return filt
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def apply_filter(
|
|
64
|
+
sinogram: torch.Tensor,
|
|
65
|
+
filter_name: FilterType = "ramp",
|
|
66
|
+
) -> torch.Tensor:
|
|
67
|
+
"""
|
|
68
|
+
Apply FBP filter to sinogram in frequency domain.
|
|
69
|
+
|
|
70
|
+
Args:
|
|
71
|
+
sinogram: Sinogram of shape [B, 1, n_angles, n_det]
|
|
72
|
+
filter_name: Type of filter to apply
|
|
73
|
+
|
|
74
|
+
Returns:
|
|
75
|
+
Filtered sinogram of shape [B, 1, n_angles, n_det]
|
|
76
|
+
"""
|
|
77
|
+
if filter_name == "none":
|
|
78
|
+
return sinogram
|
|
79
|
+
|
|
80
|
+
B, C, n_angles, n_det = sinogram.shape
|
|
81
|
+
device = sinogram.device
|
|
82
|
+
dtype = sinogram.dtype
|
|
83
|
+
|
|
84
|
+
# Pad to next power of 2 for efficient FFT (and to avoid circular conv)
|
|
85
|
+
pad_len = max(64, int(2 ** np.ceil(np.log2(2 * n_det))))
|
|
86
|
+
|
|
87
|
+
# Get filter
|
|
88
|
+
filt = get_filter(pad_len, filter_name, device=device, dtype=dtype)
|
|
89
|
+
|
|
90
|
+
# FFT of sinogram with zero-padding
|
|
91
|
+
sino_fft = fft.fft(sinogram, n=pad_len, dim=-1)
|
|
92
|
+
|
|
93
|
+
# Apply filter in frequency domain
|
|
94
|
+
filtered_fft = sino_fft * filt.view(1, 1, 1, -1)
|
|
95
|
+
|
|
96
|
+
# Inverse FFT
|
|
97
|
+
filtered = fft.ifft(filtered_fft, dim=-1).real
|
|
98
|
+
|
|
99
|
+
# Crop to original size
|
|
100
|
+
filtered = filtered[..., :n_det]
|
|
101
|
+
|
|
102
|
+
return filtered
|
torchtomo/parallel.py
ADDED
|
@@ -0,0 +1,217 @@
|
|
|
1
|
+
"""Parallel beam CT projector."""
|
|
2
|
+
|
|
3
|
+
from typing import Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
import torch.nn.functional as F
|
|
8
|
+
|
|
9
|
+
from .base import BaseProjector
|
|
10
|
+
from .filters import FilterType, apply_filter
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ParallelBeam(BaseProjector):
|
|
14
|
+
"""
|
|
15
|
+
Parallel beam CT projector.
|
|
16
|
+
|
|
17
|
+
In parallel beam geometry, all X-rays are parallel for each projection
|
|
18
|
+
angle. This is the simplest geometry and is used in synchrotron CT.
|
|
19
|
+
|
|
20
|
+
Example:
|
|
21
|
+
>>> projector = ParallelBeam(img_size=256, n_angles=180, n_det=256)
|
|
22
|
+
>>> sinogram = projector.forward(image) # [B, 1, 180, 256]
|
|
23
|
+
>>> recon = projector.fbp(sinogram) # [B, 1, 256, 256]
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
img_size: int = 256,
|
|
29
|
+
n_angles: int = 180,
|
|
30
|
+
n_det: Optional[int] = None,
|
|
31
|
+
angle_range: tuple[float, float] = (0, np.pi),
|
|
32
|
+
circle: bool = True,
|
|
33
|
+
):
|
|
34
|
+
"""
|
|
35
|
+
Initialize parallel beam projector.
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
img_size: Image size (assumed square)
|
|
39
|
+
n_angles: Number of projection angles
|
|
40
|
+
n_det: Number of detector elements (default: img_size)
|
|
41
|
+
angle_range: Range of angles in radians (default: 0 to pi)
|
|
42
|
+
circle: If True, mask image to inscribed circle
|
|
43
|
+
"""
|
|
44
|
+
n_det = n_det or img_size
|
|
45
|
+
super().__init__(img_size, n_angles, n_det, angle_range)
|
|
46
|
+
|
|
47
|
+
self.circle = circle
|
|
48
|
+
|
|
49
|
+
# Pixel size (assuming image spans [-1, 1])
|
|
50
|
+
self.pixel_size = 2.0 / img_size
|
|
51
|
+
|
|
52
|
+
# Precompute rotation grids for forward projection
|
|
53
|
+
grids = self._precompute_forward_grids()
|
|
54
|
+
self.register_buffer("forward_grids", grids)
|
|
55
|
+
|
|
56
|
+
# Circle mask for reconstruction
|
|
57
|
+
if circle:
|
|
58
|
+
coords = torch.linspace(-1, 1, img_size)
|
|
59
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
60
|
+
mask = (x**2 + y**2 <= 1).float()
|
|
61
|
+
self.register_buffer("circle_mask", mask)
|
|
62
|
+
|
|
63
|
+
def _precompute_forward_grids(self) -> torch.Tensor:
|
|
64
|
+
"""Precompute sampling grids for forward projection."""
|
|
65
|
+
grids = []
|
|
66
|
+
|
|
67
|
+
for angle in self.angles:
|
|
68
|
+
grid = self._rotation_grid(angle)
|
|
69
|
+
grids.append(grid)
|
|
70
|
+
|
|
71
|
+
return torch.stack(grids) # [n_angles, H, W, 2]
|
|
72
|
+
|
|
73
|
+
def _rotation_grid(self, angle: torch.Tensor) -> torch.Tensor:
|
|
74
|
+
"""Create sampling grid for rotating image by angle."""
|
|
75
|
+
cos_a = torch.cos(angle)
|
|
76
|
+
sin_a = torch.sin(angle)
|
|
77
|
+
|
|
78
|
+
# Create normalized coordinate grid [-1, 1]
|
|
79
|
+
coords = torch.linspace(-1, 1, self.img_size)
|
|
80
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
81
|
+
|
|
82
|
+
# Rotation matrix (rotate coordinates, not image)
|
|
83
|
+
x_rot = cos_a * x + sin_a * y
|
|
84
|
+
y_rot = -sin_a * x + cos_a * y
|
|
85
|
+
|
|
86
|
+
# Stack to grid format [H, W, 2]
|
|
87
|
+
grid = torch.stack([x_rot, y_rot], dim=-1)
|
|
88
|
+
|
|
89
|
+
return grid
|
|
90
|
+
|
|
91
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
92
|
+
"""
|
|
93
|
+
Forward projection (Radon transform): image -> sinogram.
|
|
94
|
+
|
|
95
|
+
Args:
|
|
96
|
+
x: Image tensor [B, 1, H, W]
|
|
97
|
+
|
|
98
|
+
Returns:
|
|
99
|
+
Sinogram [B, 1, n_angles, n_det]
|
|
100
|
+
"""
|
|
101
|
+
B = x.shape[0]
|
|
102
|
+
device = x.device
|
|
103
|
+
|
|
104
|
+
# Apply circle mask if enabled
|
|
105
|
+
if self.circle:
|
|
106
|
+
x = x * self.circle_mask.view(1, 1, self.img_size, self.img_size)
|
|
107
|
+
|
|
108
|
+
projections = []
|
|
109
|
+
|
|
110
|
+
for i in range(self.n_angles):
|
|
111
|
+
# Get rotation grid for this angle
|
|
112
|
+
grid = self.forward_grids[i].unsqueeze(0).expand(B, -1, -1, -1)
|
|
113
|
+
grid = grid.to(device)
|
|
114
|
+
|
|
115
|
+
# Rotate image
|
|
116
|
+
rotated = F.grid_sample(
|
|
117
|
+
x, grid, mode="bilinear", padding_mode="zeros", align_corners=True
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
# Sum along vertical axis (parallel rays) - this is the line integral
|
|
121
|
+
# Multiply by pixel size to get proper integral
|
|
122
|
+
projection = rotated.sum(dim=2) * self.pixel_size # [B, 1, W]
|
|
123
|
+
|
|
124
|
+
projections.append(projection)
|
|
125
|
+
|
|
126
|
+
# Stack to sinogram [B, 1, n_angles, n_det]
|
|
127
|
+
sinogram = torch.stack(projections, dim=2)
|
|
128
|
+
|
|
129
|
+
return sinogram
|
|
130
|
+
|
|
131
|
+
def backward(self, sinogram: torch.Tensor) -> torch.Tensor:
|
|
132
|
+
"""
|
|
133
|
+
Back projection (adjoint of Radon transform): sinogram -> image.
|
|
134
|
+
|
|
135
|
+
This smears each projection back across the image.
|
|
136
|
+
|
|
137
|
+
Args:
|
|
138
|
+
sinogram: Sinogram [B, 1, n_angles, n_det]
|
|
139
|
+
|
|
140
|
+
Returns:
|
|
141
|
+
Back-projected image [B, 1, H, W]
|
|
142
|
+
"""
|
|
143
|
+
B = sinogram.shape[0]
|
|
144
|
+
device = sinogram.device
|
|
145
|
+
|
|
146
|
+
recon = torch.zeros(B, 1, self.img_size, self.img_size, device=device)
|
|
147
|
+
|
|
148
|
+
for i in range(self.n_angles):
|
|
149
|
+
angle = self.angles[i]
|
|
150
|
+
cos_a = torch.cos(angle)
|
|
151
|
+
sin_a = torch.sin(angle)
|
|
152
|
+
|
|
153
|
+
# Image coordinates
|
|
154
|
+
coords = torch.linspace(-1, 1, self.img_size, device=device)
|
|
155
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
156
|
+
|
|
157
|
+
# Project each pixel onto detector line
|
|
158
|
+
# t = x * cos(angle) + y * sin(angle)
|
|
159
|
+
# Note: flip y to match image coordinate convention
|
|
160
|
+
t = x * cos_a - y * sin_a
|
|
161
|
+
|
|
162
|
+
# Create sampling grid for this projection
|
|
163
|
+
# We need to sample from sinogram row at position t
|
|
164
|
+
grid = torch.zeros(B, self.img_size, self.img_size, 2, device=device)
|
|
165
|
+
grid[..., 0] = t # detector position
|
|
166
|
+
grid[..., 1] = 0 # single row
|
|
167
|
+
|
|
168
|
+
# Get sinogram row [B, 1, 1, n_det]
|
|
169
|
+
sino_row = sinogram[:, :, i : i + 1, :]
|
|
170
|
+
|
|
171
|
+
# Sample and add to reconstruction
|
|
172
|
+
contribution = F.grid_sample(
|
|
173
|
+
sino_row,
|
|
174
|
+
grid,
|
|
175
|
+
mode="bilinear",
|
|
176
|
+
padding_mode="zeros",
|
|
177
|
+
align_corners=True,
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
recon += contribution
|
|
181
|
+
|
|
182
|
+
# Normalize by angular spacing (delta_theta)
|
|
183
|
+
delta_theta = (self.angle_range[1] - self.angle_range[0]) / self.n_angles
|
|
184
|
+
recon = recon * delta_theta
|
|
185
|
+
|
|
186
|
+
# Apply circle mask
|
|
187
|
+
if self.circle:
|
|
188
|
+
recon = recon * self.circle_mask.view(1, 1, self.img_size, self.img_size)
|
|
189
|
+
|
|
190
|
+
return recon
|
|
191
|
+
|
|
192
|
+
def fbp(
|
|
193
|
+
self, sinogram: torch.Tensor, filter_name: FilterType = "ramp"
|
|
194
|
+
) -> torch.Tensor:
|
|
195
|
+
"""
|
|
196
|
+
Filtered back-projection reconstruction.
|
|
197
|
+
|
|
198
|
+
Args:
|
|
199
|
+
sinogram: Sinogram [B, 1, n_angles, n_det]
|
|
200
|
+
filter_name: Filter type ('ramp', 'shepp-logan', 'cosine',
|
|
201
|
+
'hamming', 'hann')
|
|
202
|
+
|
|
203
|
+
Returns:
|
|
204
|
+
Reconstructed image [B, 1, H, W]
|
|
205
|
+
"""
|
|
206
|
+
# Apply ramp filter in frequency domain
|
|
207
|
+
filtered_sino = apply_filter(sinogram, filter_name)
|
|
208
|
+
|
|
209
|
+
# Scale the filtered sinogram by img_size / 2
|
|
210
|
+
# This compensates for the pixel_size scaling in forward projection
|
|
211
|
+
# and the discrete approximation of the FBP integral
|
|
212
|
+
filtered_sino = filtered_sino * (self.img_size / 2)
|
|
213
|
+
|
|
214
|
+
# Back-project
|
|
215
|
+
recon = self.backward(filtered_sino)
|
|
216
|
+
|
|
217
|
+
return recon
|
torchtomo/phantom.py
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
1
|
+
"""Test phantoms for CT reconstruction."""
|
|
2
|
+
|
|
3
|
+
from typing import Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def circle_phantom(
|
|
10
|
+
size: int = 256,
|
|
11
|
+
n_circles: int = 5,
|
|
12
|
+
device: Optional[torch.device] = None,
|
|
13
|
+
) -> torch.Tensor:
|
|
14
|
+
"""
|
|
15
|
+
Create a simple phantom with random circles.
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
size: Image size
|
|
19
|
+
n_circles: Number of circles
|
|
20
|
+
device: Target device
|
|
21
|
+
|
|
22
|
+
Returns:
|
|
23
|
+
Phantom image [1, 1, size, size]
|
|
24
|
+
"""
|
|
25
|
+
coords = torch.linspace(-1, 1, size, device=device)
|
|
26
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
27
|
+
|
|
28
|
+
phantom = torch.zeros(size, size, device=device)
|
|
29
|
+
|
|
30
|
+
# Background circle
|
|
31
|
+
phantom += 0.2 * ((x**2 + y**2) < 0.9).float()
|
|
32
|
+
|
|
33
|
+
# Random circles
|
|
34
|
+
torch.manual_seed(42)
|
|
35
|
+
for _ in range(n_circles):
|
|
36
|
+
cx = torch.rand(1).item() * 1.2 - 0.6
|
|
37
|
+
cy = torch.rand(1).item() * 1.2 - 0.6
|
|
38
|
+
r = torch.rand(1).item() * 0.2 + 0.05
|
|
39
|
+
intensity = torch.rand(1).item() * 0.8 + 0.2
|
|
40
|
+
|
|
41
|
+
circle = ((x - cx) ** 2 + (y - cy) ** 2) < r**2
|
|
42
|
+
phantom += intensity * circle.float()
|
|
43
|
+
|
|
44
|
+
# Clip to [0, 1]
|
|
45
|
+
phantom = phantom.clamp(0, 1)
|
|
46
|
+
|
|
47
|
+
return phantom.unsqueeze(0).unsqueeze(0)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def shepp_logan(
|
|
51
|
+
size: int = 256,
|
|
52
|
+
device: Optional[torch.device] = None,
|
|
53
|
+
modified: bool = True,
|
|
54
|
+
) -> torch.Tensor:
|
|
55
|
+
"""
|
|
56
|
+
Create the Shepp-Logan phantom.
|
|
57
|
+
|
|
58
|
+
The Shepp-Logan phantom is a standard test image for CT reconstruction
|
|
59
|
+
algorithms. It consists of ellipses simulating a human head.
|
|
60
|
+
|
|
61
|
+
Args:
|
|
62
|
+
size: Image size
|
|
63
|
+
device: Target device
|
|
64
|
+
modified: If True, use modified (higher contrast) version
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
Phantom image [1, 1, size, size]
|
|
68
|
+
"""
|
|
69
|
+
# Ellipse parameters: (intensity, a, b, x0, y0, phi)
|
|
70
|
+
# a, b = semi-axes, x0, y0 = center, phi = rotation angle
|
|
71
|
+
|
|
72
|
+
if modified:
|
|
73
|
+
# Modified Shepp-Logan with better contrast
|
|
74
|
+
ellipses = [
|
|
75
|
+
(1.0, 0.69, 0.92, 0, 0, 0), # Outer skull
|
|
76
|
+
(-0.8, 0.6624, 0.874, 0, -0.0184, 0), # Brain
|
|
77
|
+
(-0.2, 0.11, 0.31, 0.22, 0, -18), # Left ventricle
|
|
78
|
+
(-0.2, 0.16, 0.41, -0.22, 0, 18), # Right ventricle
|
|
79
|
+
(0.1, 0.21, 0.25, 0, 0.35, 0), # Top feature
|
|
80
|
+
(0.1, 0.046, 0.046, 0, 0.1, 0), # Small circle 1
|
|
81
|
+
(0.1, 0.046, 0.046, 0, -0.1, 0), # Small circle 2
|
|
82
|
+
(0.1, 0.046, 0.023, -0.08, -0.605, 0), # Bottom left
|
|
83
|
+
(0.1, 0.023, 0.023, 0, -0.606, 0), # Bottom center
|
|
84
|
+
(0.1, 0.023, 0.046, 0.06, -0.605, 0), # Bottom right
|
|
85
|
+
]
|
|
86
|
+
else:
|
|
87
|
+
# Original Shepp-Logan (low contrast)
|
|
88
|
+
ellipses = [
|
|
89
|
+
(2.0, 0.69, 0.92, 0, 0, 0),
|
|
90
|
+
(-0.98, 0.6624, 0.874, 0, -0.0184, 0),
|
|
91
|
+
(-0.02, 0.11, 0.31, 0.22, 0, -18),
|
|
92
|
+
(-0.02, 0.16, 0.41, -0.22, 0, 18),
|
|
93
|
+
(0.01, 0.21, 0.25, 0, 0.35, 0),
|
|
94
|
+
(0.01, 0.046, 0.046, 0, 0.1, 0),
|
|
95
|
+
(0.01, 0.046, 0.046, 0, -0.1, 0),
|
|
96
|
+
(0.01, 0.046, 0.023, -0.08, -0.605, 0),
|
|
97
|
+
(0.01, 0.023, 0.023, 0, -0.606, 0),
|
|
98
|
+
(0.01, 0.023, 0.046, 0.06, -0.605, 0),
|
|
99
|
+
]
|
|
100
|
+
|
|
101
|
+
# Create coordinate grid
|
|
102
|
+
coords = torch.linspace(-1, 1, size, device=device)
|
|
103
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
104
|
+
y = -y # Flip y to match standard orientation (details at bottom)
|
|
105
|
+
|
|
106
|
+
phantom = torch.zeros(size, size, device=device)
|
|
107
|
+
|
|
108
|
+
for intensity, a, b, x0, y0, phi in ellipses:
|
|
109
|
+
phi_rad = phi * np.pi / 180
|
|
110
|
+
|
|
111
|
+
# Rotate coordinates
|
|
112
|
+
cos_p = np.cos(phi_rad)
|
|
113
|
+
sin_p = np.sin(phi_rad)
|
|
114
|
+
|
|
115
|
+
x_rot = cos_p * (x - x0) + sin_p * (y - y0)
|
|
116
|
+
y_rot = -sin_p * (x - x0) + cos_p * (y - y0)
|
|
117
|
+
|
|
118
|
+
# Ellipse equation
|
|
119
|
+
inside = (x_rot / a) ** 2 + (y_rot / b) ** 2 <= 1
|
|
120
|
+
|
|
121
|
+
phantom += intensity * inside.float()
|
|
122
|
+
|
|
123
|
+
# Normalize to [0, 1]
|
|
124
|
+
phantom = (phantom - phantom.min()) / (phantom.max() - phantom.min() + 1e-8)
|
|
125
|
+
|
|
126
|
+
return phantom.unsqueeze(0).unsqueeze(0)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def forbild(
|
|
130
|
+
size: int = 256,
|
|
131
|
+
device: Optional[torch.device] = None,
|
|
132
|
+
) -> torch.Tensor:
|
|
133
|
+
"""
|
|
134
|
+
Create a simplified FORBILD head phantom.
|
|
135
|
+
|
|
136
|
+
A more challenging phantom with fine details for testing
|
|
137
|
+
resolution and artifact performance.
|
|
138
|
+
|
|
139
|
+
Args:
|
|
140
|
+
size: Image size
|
|
141
|
+
device: Target device
|
|
142
|
+
|
|
143
|
+
Returns:
|
|
144
|
+
Phantom image [1, 1, size, size]
|
|
145
|
+
"""
|
|
146
|
+
coords = torch.linspace(-1, 1, size, device=device)
|
|
147
|
+
y, x = torch.meshgrid(coords, coords, indexing="ij")
|
|
148
|
+
|
|
149
|
+
phantom = torch.zeros(size, size, device=device)
|
|
150
|
+
|
|
151
|
+
# Outer skull (ellipse)
|
|
152
|
+
skull = ((x / 0.85) ** 2 + (y / 0.95) ** 2) < 1
|
|
153
|
+
phantom += 0.2 * skull.float()
|
|
154
|
+
|
|
155
|
+
# Brain tissue
|
|
156
|
+
brain = ((x / 0.75) ** 2 + (y / 0.85) ** 2) < 1
|
|
157
|
+
phantom += 0.3 * brain.float()
|
|
158
|
+
|
|
159
|
+
# Ventricles (pair of ellipses)
|
|
160
|
+
vent_l = (((x - 0.2) / 0.08) ** 2 + ((y - 0.1) / 0.25) ** 2) < 1
|
|
161
|
+
vent_r = (((x + 0.2) / 0.08) ** 2 + ((y - 0.1) / 0.25) ** 2) < 1
|
|
162
|
+
phantom -= 0.3 * (vent_l | vent_r).float()
|
|
163
|
+
|
|
164
|
+
# High-contrast inserts (simulating lesions)
|
|
165
|
+
for i in range(5):
|
|
166
|
+
cx = 0.4 * np.cos(2 * np.pi * i / 5)
|
|
167
|
+
cy = 0.4 * np.sin(2 * np.pi * i / 5) - 0.1
|
|
168
|
+
r = 0.05
|
|
169
|
+
|
|
170
|
+
insert = ((x - cx) ** 2 + (y - cy) ** 2) < r**2
|
|
171
|
+
phantom += 0.5 * insert.float()
|
|
172
|
+
|
|
173
|
+
# Fine resolution pattern (line pairs)
|
|
174
|
+
for i, offset in enumerate([0.6, 0.65, 0.7, 0.75]):
|
|
175
|
+
width = 0.02 / (i + 1)
|
|
176
|
+
for j in range(3):
|
|
177
|
+
cx = offset
|
|
178
|
+
cy = -0.3 + j * width * 3
|
|
179
|
+
|
|
180
|
+
line = (torch.abs(x - cx) < width) & (torch.abs(y - cy) < width)
|
|
181
|
+
phantom += 0.4 * line.float()
|
|
182
|
+
|
|
183
|
+
# Normalize
|
|
184
|
+
phantom = phantom.clamp(0, 1)
|
|
185
|
+
|
|
186
|
+
return phantom.unsqueeze(0).unsqueeze(0)
|
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: torchtomo
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Differentiable CT Reconstruction in Pure PyTorch
|
|
5
|
+
Author: BIAI Lab
|
|
6
|
+
License-Expression: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/itu-biai/torchtomo
|
|
8
|
+
Project-URL: Repository, https://github.com/itu-biai/torchtomo
|
|
9
|
+
Project-URL: Issues, https://github.com/itu-biai/torchtomo/issues
|
|
10
|
+
Keywords: ct,tomography,reconstruction,pytorch,differentiable,deep-learning
|
|
11
|
+
Classifier: Development Status :: 4 - Beta
|
|
12
|
+
Classifier: Intended Audience :: Science/Research
|
|
13
|
+
Classifier: Operating System :: OS Independent
|
|
14
|
+
Classifier: Programming Language :: Python :: 3
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
18
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
19
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
20
|
+
Classifier: Topic :: Scientific/Engineering :: Image Processing
|
|
21
|
+
Classifier: Topic :: Scientific/Engineering :: Medical Science Apps.
|
|
22
|
+
Requires-Python: >=3.9
|
|
23
|
+
Description-Content-Type: text/markdown
|
|
24
|
+
License-File: LICENSE
|
|
25
|
+
Requires-Dist: torch>=1.10
|
|
26
|
+
Requires-Dist: numpy
|
|
27
|
+
Provides-Extra: test
|
|
28
|
+
Requires-Dist: pytest; extra == "test"
|
|
29
|
+
Requires-Dist: pytest-cov; extra == "test"
|
|
30
|
+
Requires-Dist: scikit-image; extra == "test"
|
|
31
|
+
Provides-Extra: dev
|
|
32
|
+
Requires-Dist: build; extra == "dev"
|
|
33
|
+
Requires-Dist: matplotlib; extra == "dev"
|
|
34
|
+
Requires-Dist: pytest; extra == "dev"
|
|
35
|
+
Requires-Dist: pytest-cov; extra == "dev"
|
|
36
|
+
Requires-Dist: ruff; extra == "dev"
|
|
37
|
+
Requires-Dist: scikit-image; extra == "dev"
|
|
38
|
+
Requires-Dist: twine; extra == "dev"
|
|
39
|
+
Dynamic: license-file
|
|
40
|
+
|
|
41
|
+
# TorchTomo
|
|
42
|
+
|
|
43
|
+
[](https://pypi.org/project/torchtomo/)
|
|
44
|
+
[](https://github.com/itu-biai/torchtomo/releases)
|
|
45
|
+
[](https://github.com/itu-biai/torchtomo/actions/workflows/test.yml)
|
|
46
|
+
[](https://github.com/itu-biai/torchtomo/blob/main/LICENSE)
|
|
47
|
+
|
|
48
|
+
Differentiable CT reconstruction primitives in pure PyTorch.
|
|
49
|
+
|
|
50
|
+
TorchTomo provides forward projection, adjoint backprojection, and filtered backprojection for parallel-beam and fan-beam geometries, with support for CPU, CUDA, and Apple Silicon (MPS).
|
|
51
|
+
|
|
52
|
+
## Features
|
|
53
|
+
|
|
54
|
+
- Pure PyTorch implementation with no custom CUDA build step
|
|
55
|
+
- Autograd-friendly operators for learned reconstruction pipelines
|
|
56
|
+
- Parallel-beam and fan-beam (flat detector) projectors
|
|
57
|
+
- Built-in FBP filters: `ramp`, `shepp-logan`, `cosine`, `hamming`, `hann`, `none`
|
|
58
|
+
- Built-in phantom generators for quick experiments
|
|
59
|
+
|
|
60
|
+
## Installation
|
|
61
|
+
|
|
62
|
+
```bash
|
|
63
|
+
pip install torchtomo
|
|
64
|
+
```
|
|
65
|
+
|
|
66
|
+
## Quick Start (Parallel Beam)
|
|
67
|
+
|
|
68
|
+
```python
|
|
69
|
+
import torch
|
|
70
|
+
from torchtomo import ParallelBeam, shepp_logan
|
|
71
|
+
|
|
72
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
73
|
+
|
|
74
|
+
phantom = shepp_logan(size=256, device=device) # [1, 1, 256, 256]
|
|
75
|
+
projector = ParallelBeam(img_size=256, n_angles=180, n_det=256).to(device)
|
|
76
|
+
|
|
77
|
+
sinogram = projector.forward(phantom) # [1, 1, 180, 256]
|
|
78
|
+
recon = projector.fbp(sinogram, filter_name="ramp") # [1, 1, 256, 256]
|
|
79
|
+
```
|
|
80
|
+
|
|
81
|
+
## Fan-Beam Example
|
|
82
|
+
|
|
83
|
+
```python
|
|
84
|
+
from torchtomo import FanBeam, shepp_logan
|
|
85
|
+
|
|
86
|
+
phantom = shepp_logan(size=256)
|
|
87
|
+
projector = FanBeam(
|
|
88
|
+
img_size=256,
|
|
89
|
+
n_angles=360,
|
|
90
|
+
n_det=400,
|
|
91
|
+
src_dist=500.0,
|
|
92
|
+
det_dist=500.0,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
sinogram = projector.forward(phantom)
|
|
96
|
+
recon = projector.fbp(sinogram, filter_name="hann")
|
|
97
|
+
```
|
|
98
|
+
|
|
99
|
+
## Differentiable Optimization Example
|
|
100
|
+
|
|
101
|
+
```python
|
|
102
|
+
import torch
|
|
103
|
+
import torch.nn.functional as F
|
|
104
|
+
from torchtomo import ParallelBeam
|
|
105
|
+
|
|
106
|
+
projector = ParallelBeam(img_size=256, n_angles=180)
|
|
107
|
+
x = torch.zeros(1, 1, 256, 256, requires_grad=True)
|
|
108
|
+
y = torch.randn(1, 1, 180, 256)
|
|
109
|
+
|
|
110
|
+
loss = F.mse_loss(projector.forward(x), y)
|
|
111
|
+
loss.backward() # gradients flow through projection operators
|
|
112
|
+
```
|
|
113
|
+
|
|
114
|
+
## API Snapshot
|
|
115
|
+
|
|
116
|
+
- `ParallelBeam(...)`
|
|
117
|
+
- `FanBeam(...)`
|
|
118
|
+
- `projector.forward(image)`
|
|
119
|
+
- `projector.backward(sinogram)`
|
|
120
|
+
- `projector.fbp(sinogram, filter_name="ramp")`
|
|
121
|
+
- `apply_filter(sinogram, filter_name=...)`
|
|
122
|
+
- `shepp_logan(size=..., device=...)`
|
|
123
|
+
- `circle_phantom(size=..., n_circles=..., device=...)`
|
|
124
|
+
- `torchtomo.phantom.forbild(size=..., device=...)`
|
|
125
|
+
|
|
126
|
+
## Tensor Shapes
|
|
127
|
+
|
|
128
|
+
- Image: `[B, 1, H, W]`
|
|
129
|
+
- Sinogram: `[B, 1, n_angles, n_det]`
|
|
130
|
+
|
|
131
|
+
## Development
|
|
132
|
+
|
|
133
|
+
```bash
|
|
134
|
+
git clone https://github.com/itu-biai/torchtomo.git
|
|
135
|
+
cd torchtomo
|
|
136
|
+
pip install -e ".[dev]"
|
|
137
|
+
```
|
|
138
|
+
|
|
139
|
+
```bash
|
|
140
|
+
make test
|
|
141
|
+
make lint
|
|
142
|
+
make build
|
|
143
|
+
```
|
|
144
|
+
|
|
145
|
+
## CI/CD
|
|
146
|
+
|
|
147
|
+
- `.github/workflows/test.yml`: Python test matrix on `push` and `pull_request`
|
|
148
|
+
- `.github/workflows/publish.yml`: release-triggered test matrix and PyPI publish step
|
|
149
|
+
|
|
150
|
+
## License
|
|
151
|
+
|
|
152
|
+
MIT
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
torchtomo/__init__.py,sha256=uCUkgFK14Bf-N23q1_NF6KKZlrTi-5I-z0KBVtsgSws,984
|
|
2
|
+
torchtomo/base.py,sha256=3dkEAxJlFyH1VXSy2IJx-jWSZr8IAa3bmKt3vlrkPQQ,2354
|
|
3
|
+
torchtomo/fanbeam.py,sha256=G9i4ifLl5EcQJy7PtL4P6Ro-mXWPV3eHHcxOMIj93Wo,11922
|
|
4
|
+
torchtomo/filters.py,sha256=lw838cO8PBwTViu5-w86-5HNWiYB846ohPKhubUtSzY,2897
|
|
5
|
+
torchtomo/parallel.py,sha256=HjiXhrB6Y3QDNkbnkBkDXcXDdSv9F0-A_NNV-IfGGjI,7008
|
|
6
|
+
torchtomo/phantom.py,sha256=R1S6Qo9JLfBh2kjfJ8oPf7R0bD2kA-TkKLxjLlBm2xc,5642
|
|
7
|
+
torchtomo-0.1.0.dist-info/licenses/LICENSE,sha256=rcEVLdBpwnZIu9UbcEraWPZRHolsdz2JOW9vLpVH_A8,532
|
|
8
|
+
torchtomo-0.1.0.dist-info/METADATA,sha256=GRnudNy-bBIWlBUApOQbu6PzOXFw_iQM-oUDbj9VTFk,4727
|
|
9
|
+
torchtomo-0.1.0.dist-info/WHEEL,sha256=YCfwYGOYMi5Jhw2fU4yNgwErybb2IX5PEwBKV4ZbdBo,91
|
|
10
|
+
torchtomo-0.1.0.dist-info/top_level.txt,sha256=X2Zh7WpPRwxOyD7lTUs_uVxOyolY4CShibZH0_CB-xM,10
|
|
11
|
+
torchtomo-0.1.0.dist-info/RECORD,,
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
This work is licensed under the Creative Commons Attribution–NonCommercial
|
|
2
|
+
4.0 International License (CC BY-NC 4.0).
|
|
3
|
+
|
|
4
|
+
You are free to use, share, and adapt this software for non-commercial
|
|
5
|
+
research and educational purposes, provided that appropriate attribution
|
|
6
|
+
is given to the authors and the source.
|
|
7
|
+
|
|
8
|
+
Commercial use by third parties is prohibited without prior written
|
|
9
|
+
permission from the authors. The authors retain all commercial rights.
|
|
10
|
+
|
|
11
|
+
To view a copy of this license, visit:
|
|
12
|
+
https://creativecommons.org/licenses/by-nc/4.0/
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
torchtomo
|