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
71 changes: 64 additions & 7 deletions antstorch/benchmark/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,11 @@
deliberate: it distinguishes these four dense-SyN-stage variants from
``'bspline_svf'`` below, which is a different transformation family
entirely (a stationary velocity field) despite the two sharing the word
"bspline".
"bspline". Each has a ``'_regadam'`` counterpart (``'gaussian_regadam'``,
``'sobolev_regadam'``, ``'dsti_regadam'``, ``'bspline_regadam'``) — the
same regularizer, but with ``syn_registration(optimizer='reg_adam')``
instead of the default ``'gradient_descent'`` (a lightweight port of
syntx's ``greedy.py`` Adam-momentum pattern; see the project doc, §40/§41).
- ``antstorch.bspline_flows.bspline_svf_registration()`` — the cubic
B-spline stationary-velocity-field model (``'bspline_svf'``/``'svf'``).
- ``antstorch.bspline_flows.gaussian_svf_registration()`` — the dense
Expand Down Expand Up @@ -64,6 +68,36 @@
"sobolev_syn": "sobolev",
"dsti_syn": "dsti",
"bspline_syn": "bspline",
# '_regadam' variants: same dense-SyN loop and regularizer as their
# '_syn' counterpart above, but with syn_registration(optimizer=
# 'reg_adam') instead of the default 'gradient_descent' -- see
# antstorch/syn/syn.py's _OPTIMIZERS docstring. optimizer and
# regularizer are independent knobs in syn_registration() itself (the
# Adam-moment quotient is computed from the raw gradient before
# _apply_regularizer runs, for whichever regularizer was requested), so
# every regularizer has a '_regadam' counterpart here, not just 'dsti'.
# 'dsti_regadam' was added first, specifically to test whether syntx's
# own advantage on 'dsti' (project doc, §34/§36/§39) comes from its
# Adam-momentum optimizer rather than from the full TVF architecture
# its own benchmark harness actually routes 'dsti' through; the other
# three were added afterward (project doc, §41) once that turned out to
# be a generic optimizer switch, to let the same question be asked of
# 'gaussian'/'sobolev'/'bspline'. Each is a separate, new arm --
# deliberately NOT changing its '_syn' counterpart's own default, so
# every prior run (§34-§39) stays reproducible as documented.
"gaussian_regadam": "gaussian",
"sobolev_regadam": "sobolev",
"dsti_regadam": "dsti",
"bspline_regadam": "bspline",
}
# model_lower values in _SYN_REGULARIZERS above that additionally force
# syn_registration(optimizer=...) rather than leaving it at 'gradient_descent'
# (or at whatever kwargs['optimizer'] the caller passed explicitly).
_SYN_OPTIMIZER_OVERRIDE = {
"gaussian_regadam": "reg_adam",
"sobolev_regadam": "reg_adam",
"dsti_regadam": "reg_adam",
"bspline_regadam": "reg_adam",
}
_BSPLINE_SVF_MODELS = ("bspline_svf", "svf")
_GAUSSIAN_SVF_MODELS = ("gaussian_svf",)
Expand Down Expand Up @@ -568,7 +602,18 @@ def evaluate_mindboggle_pair(
type_of_transform="SyNOnly", regularizer=..., initial_affine=...)``
-- the canonical affine already fit for this pair is supplied
directly, so only the fluid/B-spline regularizer differs between
them), ``'bspline_svf'``/``'svf'`` (dispatches to
them); ``'gaussian_regadam'``, ``'sobolev_regadam'``,
``'dsti_regadam'``, ``'bspline_regadam'`` (each the same dense SyN
stage and regularizer as its ``'_syn'`` counterpart above, but with
``syn_registration(optimizer='reg_adam')`` instead of the default
``'gradient_descent'`` -- separate arms, since ``optimizer`` and
``regularizer`` are independent in ``syn_registration()`` itself.
``'dsti_regadam'`` was added first, to test whether syntx's own
'dsti' advantage comes from its Adam-momentum optimizer rather than
from the full time-varying-velocity-field architecture its own
benchmark harness actually uses for that model; the other three
followed once that turned out to be a generic optimizer switch, not
a 'dsti'-specific one -- see the project doc), ``'bspline_svf'``/``'svf'`` (dispatches to
``antstorch.bspline_flows.bspline_svf_registration()`` -- a
different transformation family, a stationary velocity field, not a
SyN variant despite ``'bspline_syn'``/``'bspline_svf'`` sharing the
Expand Down Expand Up @@ -616,15 +661,22 @@ def evaluate_mindboggle_pair(
**kwargs
Model-specific overrides, forwarded to the underlying registration
call. Common ones: ``reg_iterations``, ``grad_step``, ``levels``
(all four ``_syn`` variants); ``flow_sigma``/
``total_sigma`` (gaussian_syn/sobolev_syn/dsti_syn); ``gaussian_sigma_mode``/
``conservative_smooth`` (gaussian_syn/sobolev_syn/dsti_syn -- both default
(all eight ``_syn``/``_regadam`` variants); ``flow_sigma``/
``total_sigma`` (gaussian/sobolev/dsti, both ``_syn`` and
``_regadam``); ``gaussian_sigma_mode``/
``conservative_smooth`` (gaussian/sobolev/dsti, both ``_syn`` and
``_regadam`` -- both default
to this port's own regularizer-formula conventions; pass
``gaussian_sigma_mode="voxel"``/``conservative_smooth=True`` to instead
reproduce ``syntx.syn``'s own default numbers, see
:func:`antstorch.syn.syn_registration`); ``update_field_mesh_size_at_base_level``/
:func:`antstorch.syn.syn_registration`); ``adam_betas``/``adam_eps``
(any ``_regadam`` variant -- forwarded to
:func:`antstorch.syn.syn_registration`'s own ``optimizer='reg_adam'``
moment-decay parameters; ``optimizer`` itself is set internally via
:data:`_SYN_OPTIMIZER_OVERRIDE` and is not a valid override here for
a ``_regadam`` model); ``update_field_mesh_size_at_base_level``/
``total_field_mesh_size_at_base_level``/``update_field_spline_distance``/
``total_field_spline_distance`` (bspline_syn); ``shrink_factors``/
``total_field_spline_distance`` (bspline_syn, bspline_regadam); ``shrink_factors``/
``smoothing_sigmas``/``mesh_size``/``spline_distance`` (bspline_svf);
``update_field_sigma``/``total_field_sigma``/``momentum``
(gaussian_svf); ``reg_iterations``/``syn_metric``/``syn_sampling``/
Expand Down Expand Up @@ -703,9 +755,14 @@ def evaluate_mindboggle_pair(
"bspline_enforce_stationary_boundary", "syn_metric", "neighborhood_radius",
"antisymmetric", "inverse_method", "in_loop_inverse_steps", "padding_mode",
"gaussian_sigma_mode", "conservative_smooth",
"optimizer", "adam_betas", "adam_eps",
):
if key in kwargs:
syn_kwargs[key] = kwargs[key]
if model_lower in _SYN_OPTIMIZER_OVERRIDE:
# A caller-supplied kwargs['optimizer'] (just forwarded above)
# still wins, so an explicit override is never silently dropped.
syn_kwargs.setdefault("optimizer", _SYN_OPTIMIZER_OVERRIDE[model_lower])
# No harness-level override needed here anymore: when neither
# update_field_mesh_size_at_base_level nor update_field_spline_distance
# is present in syn_kwargs, syn_registration() itself now defaults
Expand Down
75 changes: 75 additions & 0 deletions antstorch/syn/syn.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,20 @@
_SYN_TRANSFORM_TYPES = ("SyN", "SyNOnly")
_SIMILARITY_METRICS = ("mse", "lncc", "cc", "lncc2", "cc2", "mattes", "mi", "box_cc2", "dice")
_REGULARIZERS = ("gaussian", "sobolev", "dsti", "bspline")
# 'gradient_descent' (default): the existing CFL-bounded plain-gradient
# Eulerian update. 'reg_adam': ported from syntx's greedy.py
# GreedyRegistration.fit() -- per-voxel first/second Adam moments are
# accumulated across iterations (reset at the start of each pyramid level,
# matching syntx's "Reset Adam warm-up step at each scale level"), and the
# bias-corrected Adam step-direction quotient is regularized (via the same
# _apply_regularizer already used below) in place of the raw similarity
# gradient -- everything downstream (CFL bound, antisymmetric projection,
# Eulerian composition, in-loop inversion) is unchanged. See the project
# doc ("écart dsti", §34/§36/§39-40) for why this was added: to test
# whether syntx's own dsti-arm advantage comes from its Adam-momentum
# optimizer rather than from the full time-varying-velocity-field (TVF)
# architecture its benchmark harness actually routes 'dsti' through.
_OPTIMIZERS = ("gradient_descent", "reg_adam")


def _level_values(value, levels: int, name: str) -> tuple:
Expand Down Expand Up @@ -334,6 +348,9 @@ def _fit_syn_level(
num_levels: int,
gaussian_sigma_mode: str = "physical",
conservative_smooth: bool = False,
optimizer: str = "gradient_descent",
adam_betas: Tuple[float, float] = (0.9, 0.999),
adam_eps: float = 1e-8,
) -> Tuple[Tensor, Tensor, Tensor, Tensor, list]:
device, dtype = I_curr.device, I_curr.dtype
fixed_meta_t = metadata_tensors_from_dict(fixed_meta, device, dtype)
Expand All @@ -342,6 +359,23 @@ def _fit_syn_level(
boundary_mask = get_boundary_mask(fixed_meta["torch_shape"], device, dtype)
level_cfl_voxels = grad_step * math.sqrt(float(shrink_factor))

# reg_adam: per-voxel Adam first/second moments, reset at the start of
# every pyramid level (this function is called once per level) --
# mirrors syntx's greedy.py GreedyRegistration.fit() ("Reset Adam
# warm-up step at each scale level"). Kept as plain tensors (not a
# torch.optim.Optimizer over nn.Parameters) since warp_l2r/warp_r2l are
# not themselves the optimized leaves here -- the *gradient* fed into
# the existing regularizer/CFL/composition pipeline below is what
# changes between the two optimizer modes, nothing downstream of it.
use_reg_adam = optimizer == "reg_adam"
if use_reg_adam:
beta1, beta2 = adam_betas
exp_avg_l = torch.zeros_like(warp_l2r)
exp_avg_sq_l = torch.zeros_like(warp_l2r)
exp_avg_r = torch.zeros_like(warp_r2l)
exp_avg_sq_r = torch.zeros_like(warp_r2l)
adam_step = 0

# ITK's BSplineSyN doubles the update/total-field control-point mesh
# (like any TransformParametersAdaptor) from the coarsest pyramid level
# (level_index=0) to each successively finer one -- the B-spline analogue
Expand Down Expand Up @@ -395,6 +429,25 @@ def _fit_syn_level(
if grad_l is None or grad_r is None or not torch.isfinite(grad_l).all() or not torch.isfinite(grad_r).all():
raise FloatingPointError(f"non-finite SyN half-warp gradient at resolution level {level_index + 1}")

if use_reg_adam:
# Accumulate per-voxel Adam moments on the raw similarity
# gradient, then replace it with the bias-corrected step
# quotient -- exactly syntx greedy.py's exp_avg/exp_avg_sq
# update, just without its separate flow_sigma pre-smoothing
# step (this port applies the *regularizer* to the quotient
# below via the same _apply_regularizer call already used for
# 'gradient_descent', rather than introducing a second,
# differently-scoped smoothing pass).
adam_step += 1
exp_avg_l.mul_(beta1).add_(grad_l, alpha=1.0 - beta1)
exp_avg_sq_l.mul_(beta2).addcmul_(grad_l, grad_l, value=1.0 - beta2)
exp_avg_r.mul_(beta1).add_(grad_r, alpha=1.0 - beta1)
exp_avg_sq_r.mul_(beta2).addcmul_(grad_r, grad_r, value=1.0 - beta2)
bias_corr1 = 1.0 - beta1 ** adam_step
bias_corr2 = 1.0 - beta2 ** adam_step
grad_l = (exp_avg_l / bias_corr1) / ((exp_avg_sq_l / bias_corr2).sqrt().add_(adam_eps))
grad_r = (exp_avg_r / bias_corr1) / ((exp_avg_sq_r / bias_corr2).sqrt().add_(adam_eps))

grad_l = _apply_regularizer(
grad_l * boundary_mask, regularizer, flow_sigma, fixed_meta["spacing"],
mesh_size=update_mesh_size_level, domain=bspline_domain,
Expand Down Expand Up @@ -536,6 +589,9 @@ def syn_registration(
flow_sigma: float = 3.0,
total_sigma: float = 0.0,
regularizer: str = "gaussian",
optimizer: str = "gradient_descent",
adam_betas: Tuple[float, float] = (0.9, 0.999),
adam_eps: float = 1e-8,
update_field_mesh_size_at_base_level: Optional[int] = None,
total_field_mesh_size_at_base_level: int = 0,
update_field_spline_distance: Optional[Union[float, Sequence[float]]] = None,
Expand Down Expand Up @@ -605,6 +661,19 @@ def syn_registration(
``'gaussian'``/``'sobolev'``/``'dsti'``, plain Gaussian, disabled by
default).

``optimizer`` selects what the regularizer (``flow_sigma``) is applied
to each iteration: ``'gradient_descent'`` (default) applies it directly
to the raw similarity gradient, as above. ``'reg_adam'`` instead
accumulates per-voxel Adam first/second moments (``adam_betas``,
``adam_eps``) across the level's iterations -- reset at the start of
each pyramid level -- and applies the regularizer to the resulting
bias-corrected step-direction quotient instead, ported from syntx's
``greedy.py`` ``GreedyRegistration``. Everything else (the CFL bound,
antisymmetric projection, Eulerian composition) is unchanged; this
option exists to isolate whether an accuracy gap against another
library's regularizer comes from the regularizer itself or from the
Adam-momentum optimizer it happens to be paired with there.

``gaussian_sigma_mode`` and ``conservative_smooth`` tune two
``'gaussian'``/``'sobolev'``/``'dsti'`` regularizer-formula details that
this port intentionally resolves differently from ``syntx.syn``'s own
Expand Down Expand Up @@ -725,6 +794,8 @@ def syn_registration(
raise ValueError(f"syn_metric must be one of {_SIMILARITY_METRICS}")
if regularizer not in _REGULARIZERS:
raise ValueError(f"regularizer must be one of {_REGULARIZERS}")
if optimizer not in _OPTIMIZERS:
raise ValueError(f"optimizer must be one of {_OPTIMIZERS}")
if update_field_mesh_size_at_base_level is not None and update_field_mesh_size_at_base_level < 0:
raise ValueError("update_field_mesh_size_at_base_level must be >= 0")
if total_field_mesh_size_at_base_level < 0:
Expand Down Expand Up @@ -1054,6 +1125,9 @@ def syn_registration(
num_levels=num_levels,
gaussian_sigma_mode=gaussian_sigma_mode,
conservative_smooth=conservative_smooth,
optimizer=optimizer,
adam_betas=adam_betas,
adam_eps=adam_eps,
)
level_loss_history.append(history)

Expand Down Expand Up @@ -1151,6 +1225,7 @@ def syn_registration(
"flow_sigma": flow_sigma,
"total_sigma": total_sigma,
"regularizer": regularizer,
"optimizer": optimizer,
"gaussian_sigma_mode": gaussian_sigma_mode,
"conservative_smooth": conservative_smooth,
"update_field_mesh_size_at_base_level": resolved_update_field_mesh_size_at_base_level,
Expand Down
34 changes: 33 additions & 1 deletion tests/benchmark/test_evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,10 @@ def _assert_valid_success_record(rec, model):
assert len(rec["transforms"]["invtransforms"]) >= 1


@pytest.mark.parametrize("model", ["gaussian_syn", "sobolev_syn", "dsti_syn", "bspline_syn"])
@pytest.mark.parametrize("model", [
"gaussian_syn", "sobolev_syn", "dsti_syn", "bspline_syn",
"gaussian_regadam", "sobolev_regadam", "dsti_regadam", "bspline_regadam",
])
def test_evaluate_mindboggle_pair_syn_regularizers(mock_mindboggle_dataset, tmp_path, model):
pairs_csv, data_dir = mock_mindboggle_dataset
rec = evaluate_mindboggle_pair(
Expand Down Expand Up @@ -586,3 +589,32 @@ def test_evaluate_mindboggle_pair_shares_canonical_affine_across_models(mock_min
)
assert os.path.getmtime(affine_path) == mtime_after_first
assert rec1["affine_dice_sym"] == pytest.approx(rec2["affine_dice_sym"], abs=1e-9)


@pytest.mark.parametrize("syn_model,regadam_model", [
("gaussian_syn", "gaussian_regadam"),
("sobolev_syn", "sobolev_regadam"),
("dsti_syn", "dsti_regadam"),
("bspline_syn", "bspline_regadam"),
])
def test_regadam_arm_differs_from_its_syn_counterpart(mock_mindboggle_dataset, tmp_path, syn_model, regadam_model):
# Each '_regadam' arm is a separate arm (project doc, "écart dsti"
# §34-§41): same dense-SyN loop and regularizer as its '_syn'
# counterpart, but with syn_registration(optimizer='reg_adam') instead
# of the default 'gradient_descent'. optimizer and regularizer are
# independent knobs in syn_registration() itself, so this is exercised
# for all four regularizers, not just 'dsti' (where it was first added).
# This just confirms each new arm actually exercises a different update
# rule rather than silently aliasing its '_syn' counterpart.
pairs_csv, data_dir = mock_mindboggle_dataset
canonical_affine_dir = str(tmp_path / "canonical_affines")
common = dict(
pair_idx=0, device="cpu", pairs_csv=pairs_csv, data_dir=data_dir,
canonical_affine_dir=canonical_affine_dir, use_n4=False,
reg_iterations=[3, 2, 1, 1],
)
rec_syn = evaluate_mindboggle_pair(model=syn_model, **common)
rec_regadam = evaluate_mindboggle_pair(model=regadam_model, **common)
_assert_valid_success_record(rec_syn, syn_model)
_assert_valid_success_record(rec_regadam, regadam_model)
assert rec_syn["dice_sym"] != pytest.approx(rec_regadam["dice_sym"], abs=1e-9)
Loading
Loading