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
5 changes: 5 additions & 0 deletions antstorch/direct/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,3 +5,8 @@
``antstorch.utilities.cortical_thickness``. Public functions will be exported
here as the ITK-compatible port is added.
"""

from .bridge import kelly_kapowski
from .core import DiReCTResult, direct_cortical_thickness

__all__ = ["DiReCTResult", "direct_cortical_thickness", "kelly_kapowski"]
59 changes: 59 additions & 0 deletions antstorch/direct/bridge.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
"""ANTsImage interface for the tensor DiReCT implementation."""

import torch

from ..bspline_flows import ImageDomain
from ..syn.bridge import ants_image_to_tensor, tensor_to_ants_image
from ..syn.core.pipeline import auto_detect_device
from .core import direct_cortical_thickness


def kelly_kapowski(
segmentation,
gray_matter,
white_matter,
*,
iterations: int = 45,
gradient_step: float = 0.025,
smoothing_sigma: float = 1.0,
velocity_smoothing_variance: float = 1.5,
integration_points: int = 10,
thickness_prior: float = 10.0,
optimizer: str = "direct",
regularizer: str = "gaussian",
device=None,
verbose: bool = False,
):
"""ANTsImage-compatible DiReCT entry point implemented with PyTorch."""
if segmentation.dimension not in (2, 3):
raise ValueError("kelly_kapowski supports 2-D and 3-D images")
for name, image in (("gray_matter", gray_matter), ("white_matter", white_matter)):
if image.dimension != segmentation.dimension or image.shape != segmentation.shape:
raise ValueError(f"{name} must match the segmentation domain")
resolved_device = torch.device(device) if device is not None else auto_detect_device()
seg = ants_image_to_tensor(segmentation, resolved_device, normalize=False)
gray = ants_image_to_tensor(gray_matter, resolved_device, normalize=False)
white = ants_image_to_tensor(white_matter, resolved_device, normalize=False)
identity = tuple(
tuple(float(i == j) for j in range(segmentation.dimension))
for i in range(segmentation.dimension)
)
domain = ImageDomain(
tuple(int(v) for v in segmentation.shape),
tuple(float(v) for v in segmentation.spacing),
tuple(float(v) for v in segmentation.origin),
identity,
)
result = direct_cortical_thickness(
seg, gray, white, domain,
iterations=iterations,
gradient_step=gradient_step,
smoothing_sigma=smoothing_sigma,
velocity_smoothing_variance=velocity_smoothing_variance,
integration_points=integration_points,
thickness_prior=thickness_prior,
optimizer=optimizer,
regularizer=regularizer,
verbose=verbose,
)
return tensor_to_ants_image(result.thickness, segmentation)
157 changes: 157 additions & 0 deletions antstorch/direct/core.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,157 @@
"""Tensor implementation of the historical DiReCT thickness iteration."""

from dataclasses import dataclass
from typing import List, Tuple

import torch
from torch import Tensor

from ..bspline_flows import ImageDomain, compose_displacements, warp_image
from ..registration import reg_adam_direction
from ..syn.core.inverse import update_inverse_field_nd
from .forces import binary_contour, direct_force, gaussian_scalar, normalized_probability_gradient
from .regularization import regularize_velocity


@dataclass
class DiReCTResult:
thickness: Tensor
velocity: Tensor
energy_history: List[float]


def _invert(field: Tensor, initial: Tensor, domain: ImageDomain, iterations: int) -> Tensor:
# syn.core uses reversed vector components; DiReCT uses ITK x-y-z.
reversed_field = field.movedim(1, -1).flip(-1)
reversed_initial = initial.movedim(1, -1).flip(-1)
result = update_inverse_field_nd(
reversed_field,
reversed_initial,
steps=iterations,
method="fixed_point",
spacing=domain.spacing,
origin=domain.origin,
direction=domain.direction,
max_error_threshold=0.1,
mean_error_threshold=0.001,
)
return result.flip(-1).movedim(-1, 1)


@torch.no_grad()
def direct_cortical_thickness(
segmentation: Tensor,
gray_probability: Tensor,
white_probability: Tensor,
domain: ImageDomain,
*,
gray_label: int = 2,
white_label: int = 3,
iterations: int = 45,
gradient_step: float = 0.025,
integration_points: int = 10,
thickness_prior: float = 10.0,
smoothing_sigma: float = 1.0,
velocity_smoothing_variance: float = 1.5,
inverse_iterations: int = 20,
optimizer: str = "direct",
regularizer: str = "gaussian",
adam_betas: Tuple[float, float] = (0.9, 0.999),
adam_eps: float = 1e-8,
verbose: bool = False,
) -> DiReCTResult:
"""Estimate cortical thickness from a hard segmentation and GM/WM probabilities."""
expected_image_shape = (1, 1) + domain.torch_size
for name, value in (("segmentation", segmentation), ("gray_probability", gray_probability),
("white_probability", white_probability)):
if tuple(value.shape) != expected_image_shape:
raise ValueError(f"{name} must have shape {expected_image_shape}")
if optimizer not in ("direct", "reg_adam"):
raise ValueError("optimizer must be 'direct' or 'reg_adam'")
if iterations < 1 or integration_points < 1:
raise ValueError("iterations and integration_points must be positive")

dtype, device = gray_probability.dtype, gray_probability.device
gray_mask = (segmentation == gray_label).to(dtype)
white_mask = (segmentation == white_label).to(dtype)
matter_mask = ((gray_mask + white_mask) > 0).to(dtype)
matter_contour = binary_contour(matter_mask)
white_contour = binary_contour(white_mask)
active = ((segmentation != 0) & ((white_contour > 0) | (matter_contour > 0) | (gray_mask > 0))).to(dtype)

field_shape = (1, domain.dimension) + domain.torch_size
velocity = torch.zeros(field_shape, dtype=dtype, device=device)
integrated = torch.zeros_like(velocity)
thickness = torch.zeros_like(gray_probability)
energy_history: List[float] = []
adam_state = None

for outer in range(iterations):
forward_increment = torch.zeros_like(velocity)
inverse_field = torch.zeros_like(velocity)
inverse_increment = torch.zeros_like(velocity)
hit = torch.zeros_like(gray_probability)
total = torch.zeros_like(gray_probability)
energy = torch.zeros((), dtype=dtype, device=device)

for point in range(integration_points):
inverse_field = compose_displacements(inverse_field, inverse_increment, domain)
warped_white = warp_image(white_probability, inverse_field, domain, padding_mode="border")
warped_contour = warp_image(white_contour, inverse_field, domain, padding_mode="zeros")
warped_thickness = warp_image(thickness, inverse_field, domain, padding_mode="zeros")
gradient = normalized_probability_gradient(
warped_white, sigma=smoothing_sigma, spacing=domain.spacing
)
force = direct_force(
warped_white, gray_probability, gray_mask, gradient, gradient_step
)
forward_increment.add_(force)
energy.add_(((warped_white - gray_probability).abs() * gray_mask).sum())

if point == 0:
hit.copy_(white_contour)
norm = torch.linalg.vector_norm(integrated, dim=1, keepdim=True)
thickness.copy_(norm * white_contour)
total.copy_(thickness)
integrated.zero_()
else:
hit.add_(warped_contour * gray_mask)
total.add_(warped_thickness * gray_mask)

inverse_field.mul_(active)
velocity.mul_(active)
integrated.mul_(active)
inverse_increment.copy_(velocity * active)
integrated = _invert(inverse_field, integrated, domain, inverse_iterations)
inverse_field = _invert(integrated, inverse_field, domain, inverse_iterations)

smooth_hit = gaussian_scalar(hit, smoothing_sigma)
smooth_total = gaussian_scalar(total, smoothing_sigma)
estimate = torch.where(
smooth_hit > 0.001,
(smooth_total / smooth_hit.clamp_min(0.001)).clamp_min(0),
torch.zeros_like(smooth_total),
) * gray_mask
thickness.copy_(estimate)

update = forward_increment
if optimizer == "reg_adam":
adam_state, update_last = reg_adam_direction(
update.movedim(1, -1), adam_state, betas=adam_betas, eps=adam_eps
)
update = update_last.movedim(-1, 1)
velocity.add_(update)
if thickness_prior > 0:
fraction = (float(thickness_prior) / thickness.clamp_min(1e-6)).clamp(max=1.0)
velocity.mul_(torch.where(gray_mask > 0, fraction.square(), torch.ones_like(fraction)))
velocity = regularize_velocity(
velocity, mode=regularizer, variance=velocity_smoothing_variance, spacing=domain.spacing
)

gm_count = gray_mask.sum().clamp_min(1)
current_energy = float((energy / (gm_count * integration_points)).item())
energy_history.append(current_energy)
if verbose:
print(f"DiReCT iteration {outer + 1}/{iterations}: energy={current_energy:.8g}")

return DiReCTResult(thickness=thickness, velocity=velocity, energy_history=energy_history)
61 changes: 61 additions & 0 deletions antstorch/direct/forces.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
"""Contours and analytical forces used by DiReCT."""

import torch
from torch import Tensor

from ..syn.core.smoothing import separable_gaussian_filter


def binary_contour(mask: Tensor) -> Tensor:
"""Return the face-connected inner contour of a binary image."""
if mask.ndim not in (4, 5) or mask.shape[1] != 1:
raise ValueError("mask must have shape (N, 1, *spatial)")
binary = mask > 0
eroded = binary.clone()
spatial_dims = mask.ndim - 2
for axis in range(2, mask.ndim):
for offset in (-1, 1):
shifted = torch.roll(binary, shifts=offset, dims=axis)
boundary = [slice(None)] * binary.ndim
boundary[axis] = 0 if offset == 1 else -1
shifted[tuple(boundary)] = False
eroded &= shifted
return (binary & ~eroded).to(mask.dtype)


def gaussian_scalar(image: Tensor, sigma, spacing=None, sigma_mode="voxel") -> Tensor:
"""Apply the shared separable Gaussian implementation to a scalar image."""
field = image.movedim(1, -1)
return separable_gaussian_filter(
field, sigma=sigma, spacing=spacing, sigma_mode=sigma_mode
).movedim(-1, 1)


def normalized_probability_gradient(
probability: Tensor,
*,
sigma: float,
spacing,
epsilon: float = 1e-3,
) -> Tensor:
"""GradientRecursiveGaussian analogue with ITK-order vector components."""
smoothed = gaussian_scalar(probability, sigma, spacing=spacing, sigma_mode="physical")
spacing_torch = tuple(reversed(tuple(float(v) for v in spacing)))
derivatives = torch.gradient(smoothed, spacing=spacing_torch, dim=tuple(range(2, smoothed.ndim)))
gradient = torch.cat(tuple(reversed(derivatives)), dim=1)
norm = torch.linalg.vector_norm(gradient, dim=1, keepdim=True)
return torch.where(norm > epsilon, gradient / norm.clamp_min(epsilon), torch.zeros_like(gradient))


def direct_force(
warped_white_probability: Tensor,
gray_probability: Tensor,
gray_mask: Tensor,
gradient: Tensor,
gradient_step: float,
) -> Tensor:
"""Compute the historical DiReCT demons-like update force."""
delta = warped_white_probability - gray_probability
speed = -delta * gray_probability * gray_mask * float(gradient_step)
force = gradient * speed
return torch.nan_to_num(force)
48 changes: 48 additions & 0 deletions antstorch/direct/regularization.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
"""Velocity-field regularization for DiReCT."""

import math
import torch
from torch import Tensor

from ..syn.core.smoothing import (
apply_dsti_green_operator,
apply_sobolev_green_operator,
separable_gaussian_filter,
)


def stationary_boundary(field: Tensor) -> Tensor:
"""Set every spatial boundary of a channel-first field to zero."""
result = field.clone()
for axis in range(2, field.ndim):
first = [slice(None)] * field.ndim
last = [slice(None)] * field.ndim
first[axis] = 0
last[axis] = -1
result[tuple(first)] = 0
result[tuple(last)] = 0
return result


def regularize_velocity(
field: Tensor,
*,
mode: str = "gaussian",
variance: float = 1.5,
spacing=None,
) -> Tensor:
"""Regularize a channel-first physical displacement field."""
if variance <= 0 or mode == "none":
return stationary_boundary(field)
channel_last = field.movedim(1, -1)
if mode == "gaussian":
smoothed = separable_gaussian_filter(channel_last, sigma=math.sqrt(variance))
elif mode == "sobolev":
smoothed = apply_sobolev_green_operator(
channel_last, fluid_sigma=variance, alpha=variance, spacing=spacing
)
elif mode == "dsti":
smoothed = apply_dsti_green_operator(channel_last, fluid_sigma=variance, alpha=variance)
else:
raise ValueError("regularizer must be 'gaussian', 'sobolev', 'dsti', or 'none'")
return stationary_boundary(smoothed.movedim(-1, 1))
1 change: 1 addition & 0 deletions antstorch/utilities/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from .preprocess_image import preprocess_brain_image
from .deep_atropos import deep_atropos
from .cortical_thickness import cortical_thickness
from .cortical_thickness import cortical_thickness2
from .cortical_thickness import longitudinal_cortical_thickness
from .deep_flash import deep_flash
from .harvard_oxford_atlas_labeling import harvard_oxford_atlas_labeling
Expand Down
46 changes: 46 additions & 0 deletions antstorch/utilities/cortical_thickness.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,52 @@ def cortical_thickness(t1, device=None, verbose: bool = False):
}


def cortical_thickness2(
t1,
device=None,
optimizer: str = "direct",
regularizer: str = "gaussian",
verbose: bool = False,
):
"""Cortical-thickness workflow using ANTsTorch's tensor DiReCT engine.

This provisional entry point mirrors :func:`cortical_thickness` while
leaving its established ``ants.kelly_kapowski`` behavior unchanged.
"""
from ..direct import kelly_kapowski
from ..utilities.deep_atropos import deep_atropos

atropos = deep_atropos(
[t1, None, None], do_preprocessing=True, device=device, verbose=verbose
)
kk_segmentation = ants.image_clone(atropos["segmentation_image"])
kk_segmentation[kk_segmentation == 4] = 3
gray_matter = atropos["probability_images"][2]
white_matter = atropos["probability_images"][3] + atropos["probability_images"][4]
thickness = kelly_kapowski(
kk_segmentation,
gray_matter,
white_matter,
iterations=45,
gradient_step=0.025,
velocity_smoothing_variance=1.5,
optimizer=optimizer,
regularizer=regularizer,
device=device,
verbose=verbose,
)
return {
"thickness_image": thickness,
"segmentation_image": atropos["segmentation_image"],
"csf_probability_image": atropos["probability_images"][1],
"gray_matter_probability_image": atropos["probability_images"][2],
"white_matter_probability_image": atropos["probability_images"][3],
"deep_gray_matter_probability_image": atropos["probability_images"][4],
"brain_stem_probability_image": atropos["probability_images"][5],
"cerebellum_probability_image": atropos["probability_images"][6],
}


def longitudinal_cortical_thickness(
t1s,
initial_template: "str|ants.ANTsImage" = "oasis",
Expand Down
Loading
Loading