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 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
+ [![PyPI](https://img.shields.io/pypi/v/torchtomo.svg)](https://pypi.org/project/torchtomo/)
44
+ [![Changelog](https://img.shields.io/github/v/release/itu-biai/torchtomo?include_prereleases&label=changelog)](https://github.com/itu-biai/torchtomo/releases)
45
+ [![Tests](https://github.com/itu-biai/torchtomo/actions/workflows/test.yml/badge.svg)](https://github.com/itu-biai/torchtomo/actions/workflows/test.yml)
46
+ [![License](https://img.shields.io/badge/license-MIT-blue.svg)](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,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (82.0.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -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