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
53 changes: 48 additions & 5 deletions antstorch/bspline_flows/n4_bias_field_correction.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

import warnings
from math import ceil, log2
from time import perf_counter
from typing import Optional, Union

import torch
Expand All @@ -31,6 +32,14 @@
DEFAULT_N4_SPLINE_DISTANCE_MM = 200.0


def _synchronize_device(device: torch.device) -> None:
"""Synchronize accelerator work for accurate verbose timings."""
if device.type == "cuda":
torch.cuda.synchronize(device)
elif device.type == "mps":
torch.mps.synchronize()


def _expand_like_image(value: Optional[Tensor], image: Tensor, name: str, default: float) -> Tensor:
if value is None:
return image.new_full((image.shape[0], 1) + image.shape[2:], default).expand_as(image)
Expand Down Expand Up @@ -261,6 +270,7 @@ def n4_bias_field_correction(
bias_field_fwhm: float = 0.15,
eps: float = 1e-6,
stable_accumulation: Optional[bool] = None,
verbose: bool = False,
) -> Tensor:
"""Differentiable N4-style correction for batched 2-D/3-D scalar images.

Expand All @@ -271,6 +281,8 @@ def n4_bias_field_correction(
``spline_param`` is left as ``None`` (the default), a physical spline
distance of ``DEFAULT_N4_SPLINE_DISTANCE_MM`` (200 mm) is used,
matching real ANTs' own ``N4BiasFieldCorrection`` default.
Set ``verbose=True`` to report the convergence measurement at every
fitting iteration.
"""
if image.ndim not in (4, 5) or not image.is_floating_point():
raise ValueError("image must be a floating (N,C,H,W) or (N,C,D,H,W) tensor")
Expand Down Expand Up @@ -342,6 +354,25 @@ def n4_bias_field_correction(
)

for level, maximum_iterations in enumerate(iterations):
next_lattice_itk = (
lattice_itk
if level == 0
else tuple(2 * value - 3 for value in lattice_itk)
)
if verbose:
accumulation_mode = "stable" if stable_accumulation else "fast"
print(
f"ANTsTorch N4 preparing level {level + 1}/{len(iterations)}: "
f"lattice={next_lattice_itk}, accumulation={accumulation_mode}",
flush=True,
)
_synchronize_device(image.device)
preparation_start = perf_counter()
if level > 0:
accumulated_coefficients = refine_bspline_coefficients(
accumulated_coefficients
)
lattice_itk = next_lattice_itk
active = torch.ones(
(image.shape[0], image.shape[1]) + (1,) * dimension,
dtype=image.dtype,
Expand All @@ -354,7 +385,14 @@ def n4_bias_field_correction(
# being rebuilt on every one of ``maximum_iterations`` iterations.
geometry = _bspline_fit_geometry(shrunk_domain.torch_size, lattice_itk, image.dtype, image.device, eps)
fit_context = _bspline_fit_context(weight_flat, geometry, stable_accumulation)
for _ in range(maximum_iterations):
if verbose:
_synchronize_device(image.device)
print(
f"ANTsTorch N4 prepared level {level + 1}/{len(iterations)} "
f"in {perf_counter() - preparation_start:.3f} s",
flush=True,
)
for iteration in range(maximum_iterations):
uncorrected = log_input - log_bias
sharpened = _histogram_sharpen(
uncorrected,
Expand All @@ -380,15 +418,20 @@ def n4_bias_field_correction(
dim=tuple(range(2, image.ndim)), keepdim=True
) / (count - 1.0)
convergence_measurement = variance.sqrt() / mean.clamp_min(eps)
if verbose:
maximum_convergence = float(
convergence_measurement.detach().max().cpu()
)
print(
f"ANTsTorch N4 level {level + 1}/{len(iterations)}, "
f"iteration {iteration + 1}/{maximum_iterations}: "
f"convergence={maximum_convergence:.6g}"
)
log_bias = new_log_bias
accumulated_coefficients = accumulated_coefficients + active * coefficients
# Tensor gating honors convergence independently for every batch
# item/channel without a device-to-host synchronization.
active = active * (convergence_measurement > tolerance).to(image.dtype)
if level + 1 < len(iterations):
accumulated_coefficients = refine_bspline_coefficients(accumulated_coefficients)
lattice_itk = tuple(2 * value - 3 for value in lattice_itk)

full_log_bias = synthesize_bspline_velocity(accumulated_coefficients, domain)
bias = torch.exp(full_log_bias)
corrected = image / bias.clamp_min(eps)
Expand Down
50 changes: 48 additions & 2 deletions tools/benchmarks/compare_n4_bias_field_correction.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,10 @@
Use another 2-D or 3-D image and CUDA, if available::

python tools/benchmarks/compare_n4_bias_field_correction.py image.nii.gz --device cuda

Write the output images to a dedicated directory::

python tools/benchmarks/compare_n4_bias_field_correction.py --output-dir results/n4
"""

import argparse
Expand Down Expand Up @@ -77,13 +81,34 @@ def parse_args() -> argparse.Namespace:
help="Iterations at each fitting level (default: 20 20)",
)
parser.add_argument("--tolerance", type=float, default=0.0)
parser.add_argument(
"--verbose",
action="store_true",
help="Print iteration progress from both ANTs and ANTsTorch N4.",
)
parser.add_argument(
"--stable-accumulation",
action=argparse.BooleanOptionalAction,
default=None,
help=(
"Control deterministic ANTsTorch reductions. The default uses "
"stable accumulation on MPS and fast accumulation elsewhere; "
"use --no-stable-accumulation to avoid slow MPS level setup."
),
)
parser.add_argument(
"--mesh-size",
type=int,
nargs="+",
default=None,
help="B-spline mesh size in ITK x-y-z order (default: one span per axis)",
)
parser.add_argument(
"--output-dir",
type=Path,
default=Path("."),
help="Directory for output images; it is created if needed (default: current directory).",
)
parser.add_argument("--output-prefix", default="n4_comparison")
return parser.parse_args()

Expand All @@ -103,22 +128,31 @@ def main() -> None:
if len(mesh_size) != t1.dimension:
raise ValueError(f"--mesh-size needs {t1.dimension} values for this image")
convergence = {"iters": args.iterations, "tol": args.tolerance}
stable_accumulation = args.stable_accumulation
if stable_accumulation is None:
stable_accumulation = device.type == "mps"

start = time.perf_counter()
if args.verbose:
print("Running ANTs N4 corrected-image pass...")
n4_ants = ants.n4_bias_field_correction(
t1,
mask=mask,
shrink_factor=args.shrink_factor,
convergence=convergence,
spline_param=mesh_size,
verbose=args.verbose,
)
if args.verbose:
print("Running ANTs N4 bias-field pass...")
bias_ants = ants.n4_bias_field_correction(
t1,
mask=mask,
shrink_factor=args.shrink_factor,
convergence=convergence,
spline_param=mesh_size,
return_bias_field=True,
verbose=args.verbose,
)
ants_seconds = time.perf_counter() - start

Expand All @@ -135,14 +169,20 @@ def main() -> None:
torch.cuda.reset_peak_memory_stats(device)
synchronize(device)
start = time.perf_counter()
if args.verbose:
print("Running ANTsTorch N4 corrected-image pass...")
n4_torch_tensor = antstorch.n4_bias_field_correction(
t1_tensor,
domain,
mask_tensor,
shrink_factor=args.shrink_factor,
convergence=convergence,
spline_param=tuple(mesh_size),
stable_accumulation=stable_accumulation,
verbose=args.verbose,
)
if args.verbose:
print("Running ANTsTorch N4 bias-field pass...")
bias_torch_tensor = antstorch.n4_bias_field_correction(
t1_tensor,
domain,
Expand All @@ -151,6 +191,8 @@ def main() -> None:
convergence=convergence,
spline_param=tuple(mesh_size),
return_bias_field=True,
stable_accumulation=stable_accumulation,
verbose=args.verbose,
)
synchronize(device)
torch_seconds = time.perf_counter() - start
Expand All @@ -166,7 +208,8 @@ def main() -> None:
normalized_torch_bias = normalized_bias_array(bias_torch)
bias_difference = np.log(normalized_torch_bias) - np.log(normalized_ants_bias)

prefix = Path(args.output_prefix)
args.output_dir.mkdir(parents=True, exist_ok=True)
prefix = args.output_dir / args.output_prefix
ants.image_write(n4_ants, f"{prefix}_ants_corrected.nii.gz")
ants.image_write(n4_torch, f"{prefix}_antstorch_corrected.nii.gz")
ants.image_write(bias_ants, f"{prefix}_ants_bias.nii.gz")
Expand All @@ -180,7 +223,10 @@ def main() -> None:
print(f"ANTsTorch intensity range: {corrected_torch_array.min():.6g} to {corrected_torch_array.max():.6g}")
print(f"ANTs bias-field range: {normalized_ants_bias.min():.6g} to {normalized_ants_bias.max():.6g}")
print(f"ANTsTorch bias-field range: {normalized_torch_bias.min():.6g} to {normalized_torch_bias.max():.6g}")
print(f"B-spline accumulation: {'stable matrix reduction' if device.type == 'mps' else 'vectorized scatter'}")
print(
"B-spline accumulation: "
f"{'stable matrix reduction' if stable_accumulation else 'vectorized scatter'}"
)
print(f"Corrected-image RMSE: {np.sqrt(np.mean(corrected_difference**2)):.6g}")
print(
"Scale-aligned corrected-image RMSE: "
Expand Down
Loading