Skip to content
Open
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
20 changes: 10 additions & 10 deletions jitfields/_regularisers_fields.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -336,7 +336,7 @@ def field_kernel_add(
out: Optional[Tensor] = None,
_sub: bool = False,
) -> Tensor:
"""See `flow_kernel`"""
"""See `field_kernel`"""
# allocate output
if out is None:
out = inp.clone()
Expand All@@ -352,9 +352,9 @@ def field_kernel_add(

# forward
impl = cuda_impl if out.is_cuda else cpu_impl
impl.flow_kernel(out, bound, voxel_size,
absolute, membrane, bending,
'sub' if _sub else 'add')
impl.field_kernel(out, bound, voxel_size,
absolute, membrane, bending,
'sub' if _sub else 'add')

return out

Expand All@@ -369,7 +369,7 @@ def field_kernel_add_(
voxel_size: OneOrSeveral[float] = 1,
_sub: bool = False,
) -> Tensor:
"""See `flow_kernel`"""
"""See `field_kernel`"""
nc = inp.shape[-1]
bound = ensure_list(bound, ndim)
voxel_size = make_vector(voxel_size, ndim).tolist()
Expand All@@ -379,9 +379,9 @@ def field_kernel_add_(

# forward
impl = cuda_impl if inp.is_cuda else cpu_impl
impl.flow_kernel(inp, bound, voxel_size,
absolute, membrane, bending,
'sub' if _sub else 'add')
impl.field_kernel(inp, bound, voxel_size,
absolute, membrane, bending,
'sub' if _sub else 'add')
return inp


Expand All@@ -397,7 +397,7 @@ def field_kernel_sub(
) -> Tensor:
"""See `field_kernel`"""
return field_kernel_add(ndim, inp, absolute, membrane, bending,
bound, voxel_size, out)
bound, voxel_size, out, True)


def field_kernel_sub_(
Expand All@@ -411,7 +411,7 @@ def field_kernel_sub_(
) -> Tensor:
"""See `field_kernel`"""
return field_kernel_add_(ndim, inp, absolute, membrane, bending,
bound, voxel_size)
bound, voxel_size, True)


def field_diag(
Expand Down
4 changes: 2 additions & 2 deletions jitfields/_regularisers_flows.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -425,7 +425,7 @@ def flow_kernel_sub(
from `inp`.
"""
return flow_kernel_add(inp, absolute, membrane, bending, shears, div,
bound, voxel_size, out)
bound, voxel_size, out, True)


def flow_kernel_sub_(
Expand All@@ -443,7 +443,7 @@ def flow_kernel_sub_(
from `inp` (inplace).
"""
return flow_kernel_add_(inp, absolute, membrane, bending, shears, div,
bound, voxel_size)
bound, voxel_size, True)


def flow_diag(
Expand Down
45 changes: 44 additions & 1 deletion jitfields/tests/test_reg_field.py
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,8 @@
from jitfields.regularisers import field_kernel, field_matvec
from jitfields.regularisers import (
field_kernel, field_matvec,
field_kernel_add, field_kernel_add_,
field_kernel_sub, field_kernel_sub_,
)
from .ref_kernels import kernels1, kernels2, kernels3
from .utils import get_test_devices, init_device
import torch
Expand DownExpand Up@@ -44,6 +48,45 @@ def test_kernel(device, dim):
f"{pred_bending}\n{kernel_bending}"


@pytest.mark.parametrize("device", devices)
@pytest.mark.parametrize("dim", dims)
def test_kernel_add_sub(device, dim):
"""`field_kernel_{add,sub}[_]` must agree with `field_kernel`.

The voxel size is deliberately anisotropic: under the default
(isotropic) voxel size the field and flow kernels happen to be
numerically identical, so a variant that computes the *flow* kernel
instead of the *field* kernel would go unnoticed.
"""
device = init_device(device)
backend = dict(device=device, dtype=torch.float32)

voxel_size = [1.5, 1., 2.5][:dim]
absolute, membrane, bending = 0.3, 1., 0.2

kernel = field_kernel([5]*dim, absolute=absolute, membrane=membrane,
bending=bending, voxel_size=voxel_size, **backend)
inp = torch.randn(kernel.shape, **backend)

add, sub = inp + kernel, inp - kernel

pred = field_kernel_add(dim, inp, absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, add), f"{pred}\n{add}"

pred = field_kernel_sub(dim, inp, absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, sub), f"{pred}\n{sub}"

pred = field_kernel_add_(dim, inp.clone(), absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, add), f"{pred}\n{add}"

pred = field_kernel_sub_(dim, inp.clone(), absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, sub), f"{pred}\n{sub}"


@pytest.mark.parametrize("device", devices)
@pytest.mark.parametrize("dim", dims)
def test_matvec(device, dim):
Expand Down
45 changes: 44 additions & 1 deletion jitfields/tests/test_reg_flow.py
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,8 @@
from jitfields.regularisers import flow_kernel, flow_matvec
from jitfields.regularisers import (
flow_kernel, flow_matvec,
flow_kernel_add, flow_kernel_add_,
flow_kernel_sub, flow_kernel_sub_,
)
from .ref_kernels import kernels1, kernels2, kernels3
from .utils import get_test_devices, init_device
import torch
Expand DownExpand Up@@ -55,6 +59,45 @@ def test_kernel(device, dim):
assert torch.allclose(pred_div, kernel_div), f"{pred_div}\n{kernel_div}"


@pytest.mark.parametrize("device", devices)
@pytest.mark.parametrize("dim", dims)
def test_kernel_add_sub(device, dim):
"""`flow_kernel_{add,sub}[_]` must agree with `flow_kernel`.

The voxel size is deliberately anisotropic: under the default
(isotropic) voxel size the field and flow kernels happen to be
numerically identical, so a variant that computes the wrong kernel
would go unnoticed.
"""
device = init_device(device)
backend = dict(device=device, dtype=torch.float32)

voxel_size = [1.5, 1., 2.5][:dim]
absolute, membrane, bending = 0.3, 1., 0.2

kernel = flow_kernel([5]*dim, absolute=absolute, membrane=membrane,
bending=bending, voxel_size=voxel_size, **backend)
inp = torch.randn(kernel.shape, **backend)

add, sub = inp + kernel, inp - kernel

pred = flow_kernel_add(inp, absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, add), f"{pred}\n{add}"

pred = flow_kernel_sub(inp, absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, sub), f"{pred}\n{sub}"

pred = flow_kernel_add_(inp.clone(), absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, add), f"{pred}\n{add}"

pred = flow_kernel_sub_(inp.clone(), absolute, membrane, bending,
voxel_size=voxel_size)
assert torch.allclose(pred, sub), f"{pred}\n{sub}"


@pytest.mark.parametrize("device", devices)
@pytest.mark.parametrize("dim", dims)
def test_matvec(device, dim):
Expand Down