Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions antstorch/benchmark/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,7 +214,7 @@ def get_n4_cached_subject_volume(
) -> ants.ANTsImage:
"""Loads an N4-bias-corrected subject volume from disk cache, or computes and caches it.

Uses :func:`antstorch.bspline_flows.n4_bias_field_correction` directly
Uses :func:`antstorch.bspline_flows.n4_bias_field_correction_tensor` directly
(an in-package call, unlike ``syntx.benchmark.data``'s own version of
this function, which reaches ``antstorch`` as an external, optional
dependency).
Expand All @@ -232,7 +232,7 @@ def get_n4_cached_subject_volume(
raw_img = ants.image_read(raw_brain_path)
try:
import torch
from antstorch.bspline_flows import n4_bias_field_correction
from antstorch.bspline_flows import n4_bias_field_correction_tensor

arr = raw_img.numpy()
tensor = torch.from_numpy(arr.transpose(2, 1, 0)).unsqueeze(0).unsqueeze(0).float()
Expand All @@ -243,7 +243,7 @@ def get_n4_cached_subject_volume(
if verbose:
print(f"[antstorch.benchmark] Computing N4 correction for {subject}...", flush=True)

corrected_tensor = n4_bias_field_correction(
corrected_tensor = n4_bias_field_correction_tensor(
tensor,
mask=mask,
shrink_factor=4,
Expand Down
4 changes: 2 additions & 2 deletions antstorch/bspline_flows/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
synthesize_bspline_velocity,
)
from .deterministic_registration import DeterministicBSplineRegistration
from .n4_bias_field_correction import DEFAULT_N4_SPLINE_DISTANCE_MM, N4BiasFieldCorrection, n4_bias_field_correction
from .n4_bias_field_correction import DEFAULT_N4_SPLINE_DISTANCE_MM, N4BiasFieldCorrection, n4_bias_field_correction_tensor
from .physical_gradient_descent import PhysicalGradientDescent
from .gaussian_svf_registration import gaussian_svf_registration
from .bspline_svf_registration import DEFAULT_BSPLINE_SPLINE_DISTANCE_MM, bspline_svf_registration
Expand Down Expand Up @@ -49,7 +49,7 @@
"synthesize_bspline_velocity",
"DeterministicBSplineRegistration",
"N4BiasFieldCorrection",
"n4_bias_field_correction",
"n4_bias_field_correction_tensor",
"DEFAULT_N4_SPLINE_DISTANCE_MM",
"PhysicalGradientDescent",
"gaussian_svf_registration",
Expand Down
2 changes: 1 addition & 1 deletion antstorch/bspline_flows/bspline_scattered_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
coefficient lattice -> exactly refine that lattice for the next, finer
level (see :func:`~.bspline_synthesis.refine_bspline_coefficients`). This is
the same accumulate/refine pattern
:func:`~.n4_bias_field_correction.n4_bias_field_correction` uses internally,
:func:`~.n4_bias_field_correction.n4_bias_field_correction_tensor` uses internally,
generalized here from N4's regular shrunk-image grid to arbitrary scattered
points with independent parametric locations.
"""
Expand Down
4 changes: 2 additions & 2 deletions antstorch/bspline_flows/n4_bias_field_correction.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,7 +254,7 @@ def _initial_lattice_size(domain: ImageDomain, spline_param) -> tuple:
return tuple(int(value) + 3 for value in values)


def n4_bias_field_correction(
def n4_bias_field_correction_tensor(
image: Tensor,
domain: Optional[ImageDomain] = None,
mask: Optional[Tensor] = None,
Expand Down Expand Up @@ -454,4 +454,4 @@ def __init__(self, **kwargs):
self.kwargs = kwargs

def forward(self, image: Tensor, domain: Optional[ImageDomain] = None, mask: Optional[Tensor] = None, weight_mask: Optional[Tensor] = None) -> Tensor:
return n4_bias_field_correction(image, domain, mask, weight_mask=weight_mask, **self.kwargs)
return n4_bias_field_correction_tensor(image, domain, mask, weight_mask=weight_mask, **self.kwargs)
1 change: 1 addition & 0 deletions antstorch/utilities/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from .cortical_thickness import cortical_thickness
from .cortical_thickness import cortical_thickness2
from .cortical_thickness import longitudinal_cortical_thickness
from .n4_bias_field_correction import n4_bias_field_correction
from .deep_flash import deep_flash
from .harvard_oxford_atlas_labeling import harvard_oxford_atlas_labeling
from .desikan_killiany_tourville_labeling import desikan_killiany_tourville_labeling
Expand Down
137 changes: 137 additions & 0 deletions antstorch/utilities/n4_bias_field_correction.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
"""ANTsImage interface for differentiable ANTsTorch N4 correction."""

from typing import Optional

import ants
import torch

from ..bspline_flows import ImageDomain
from ..bspline_flows import n4_bias_field_correction_tensor
from ..syn.bridge import ants_image_to_tensor, tensor_to_ants_image
from .device_manager import get_default_device


def _validate_scalar_image(image, name: str) -> None:
if not ants.is_image(image):
raise TypeError(f"{name} must be an ANTsImage")
if image.dimension not in (2, 3) or image.components != 1:
raise ValueError(f"{name} must be a scalar 2-D or 3-D ANTsImage")


def _validate_optional_image(image, reference, name: str) -> None:
if image is None:
return
_validate_scalar_image(image, name)
if image.shape != reference.shape or not ants.image_physical_space_consistency(
reference, image
):
raise ValueError(f"{name} must occupy the same physical space as image")


def n4_bias_field_correction(
image,
mask=None,
*,
rescale_intensities: bool = False,
shrink_factor: int = 4,
convergence: Optional[dict] = None,
spline_param=None,
return_bias_field: bool = False,
weight_mask=None,
number_of_histogram_bins: int = 200,
wiener_filter_noise: float = 0.01,
bias_field_fwhm: float = 0.15,
stable_accumulation: Optional[bool] = None,
device=None,
verbose: bool = False,
):
"""Correct an ANTsImage with ANTsTorch's differentiable N4 engine.

This provisional high-level interface mirrors the principal options of
:func:`ants.n4_bias_field_correction`. It converts ANTs images to tensors,
runs the tensor implementation, and restores the input image geometry.

Parameters
----------
image : ANTsImage
Scalar 2-D or 3-D image to correct.
mask : ANTsImage, optional
Nonzero voxels define the correction domain. The default includes the
full image.
rescale_intensities : bool
Rescale corrected intensities to the masked input range.
shrink_factor : int
Subsampling factor used while estimating the bias field.
convergence : dict, optional
``{"iters": [...], "tol": value}`` fitting schedule.
spline_param : float or sequence, optional
Physical knot spacing when scalar, or B-spline mesh size in ITK
x-y-z order when a sequence.
return_bias_field : bool
Return the multiplicative bias field instead of the corrected image.
weight_mask : ANTsImage, optional
Nonnegative confidence weights in the same physical space as ``image``.
number_of_histogram_bins : int
Number of bins used for histogram sharpening.
wiener_filter_noise : float
Wiener filter noise parameter.
bias_field_fwhm : float
Bias-field full width at half maximum.
stable_accumulation : bool, optional
Use deterministic matrix reductions. The tensor engine defaults to
this mode on MPS and to vectorized scatter elsewhere.
device : str or torch.device, optional
PyTorch device. The configured ANTsTorch default is used when omitted.
verbose : bool
Report per-level setup time and per-iteration convergence.

Returns
-------
ANTsImage
Corrected image, or the bias field when ``return_bias_field=True``.
"""
_validate_scalar_image(image, "image")
_validate_optional_image(mask, image, "mask")
_validate_optional_image(weight_mask, image, "weight_mask")

resolved_device = (
torch.device(device) if device is not None else get_default_device()
)
domain = ImageDomain(
size=tuple(int(value) for value in image.shape),
spacing=tuple(float(value) for value in image.spacing),
origin=tuple(float(value) for value in image.origin),
direction=tuple(
tuple(float(value) for value in row) for row in image.direction
),
)
image_tensor = ants_image_to_tensor(
image, resolved_device, normalize=False
)
mask_tensor = (
ants_image_to_tensor(mask, resolved_device, normalize=False)
if mask is not None
else None
)
weight_tensor = (
ants_image_to_tensor(weight_mask, resolved_device, normalize=False)
if weight_mask is not None
else None
)
result = n4_bias_field_correction_tensor(
image_tensor,
domain,
mask_tensor,
rescale_intensities=rescale_intensities,
shrink_factor=shrink_factor,
convergence=convergence,
spline_param=spline_param,
return_bias_field=return_bias_field,
weight_mask=weight_tensor,
number_of_histogram_bins=number_of_histogram_bins,
wiener_filter_noise=wiener_filter_noise,
bias_field_fwhm=bias_field_fwhm,
stable_accumulation=stable_accumulation,
verbose=verbose,
)
return tensor_to_ants_image(result, image)
8 changes: 4 additions & 4 deletions docs/antsx_tutorial_bspline_flows.md
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@ The ANTsTorch implementation accepts batched 2-D and 3-D tensors. The mask
may contain one channel or the same number of channels as the image.

```python
from antstorch.bspline_flows import n4_bias_field_correction
from antstorch.bspline_flows import n4_bias_field_correction_tensor

device = (
"cuda"
Expand All @@ -90,7 +90,7 @@ r16_domain = ants_domain(r16)

convergence = {"iters": [50, 50, 50, 50], "tol": 1e-7}

r16_n4_tensor = n4_bias_field_correction(
r16_n4_tensor = n4_bias_field_correction_tensor(
r16_tensor,
domain=r16_domain,
mask=r16_mask_tensor,
Expand All @@ -100,7 +100,7 @@ r16_n4_tensor = n4_bias_field_correction(
rescale_intensities=True,
)

r16_bias_tensor = n4_bias_field_correction(
r16_bias_tensor = n4_bias_field_correction_tensor(
r16_tensor,
domain=r16_domain,
mask=r16_mask_tensor,
Expand All @@ -122,7 +122,7 @@ The tensors remain differentiable with respect to the input image:

```python
differentiable_input = r16_tensor.detach().clone().requires_grad_(True)
corrected = n4_bias_field_correction(
corrected = n4_bias_field_correction_tensor(
differentiable_input,
domain=r16_domain,
mask=r16_mask_tensor,
Expand Down
Loading
Loading