from __future__ import annotations
from itertools import product
from typing import Literal
import torch
from torch_lattice import SparseTensor
from torch_lattice.utils import make_ntuple
from .relation import (
build_pool_output_coords,
build_target_out_in_map,
build_target_transposed_out_in_map,
build_transposed_output_coords,
)
__all__ = [
"avg_pool3d",
"global_pool",
"global_sum_pool",
"global_avg_pool",
"global_max_pool",
"max_pool3d",
"pool3d",
"pool_transpose3d",
"sum_pool3d",
"trilinear_upsample3d",
]
PoolMode = Literal["sum", "max", "avg"]
[docs]
def pool3d(
inputs: SparseTensor,
*,
mode: PoolMode,
kernel_size=2,
stride=2,
padding=0,
dilation=1,
) -> SparseTensor:
"""Local sparse 3D pooling over convolution-style neighborhoods."""
if mode not in {"sum", "max", "avg"}:
raise ValueError("pool3d mode must be 'sum', 'max', or 'avg'.")
kernel_size = make_ntuple(kernel_size, ndim=3)
stride = make_ntuple(stride, ndim=3)
padding = make_ntuple(padding, ndim=3)
dilation = make_ntuple(dilation, ndim=3)
output_coords = build_pool_output_coords(
inputs.coords,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
spatial_range=inputs.spatial_range,
)
relation = build_target_out_in_map(
inputs.coords,
output_coords,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
)
feats = _pool_features(inputs.feats, relation, mode)
output_stride = tuple(inputs.stride[index] * stride[index] for index in range(3))
return inputs.with_coordinates(
feats=feats,
coords=output_coords,
stride=output_stride,
spatial_range=_pooled_spatial_range(
inputs.spatial_range,
kernel_size,
stride,
padding,
dilation,
),
)
[docs]
def sum_pool3d(inputs: SparseTensor, **kwargs) -> SparseTensor:
return pool3d(inputs, mode="sum", **kwargs)
[docs]
def max_pool3d(inputs: SparseTensor, **kwargs) -> SparseTensor:
return pool3d(inputs, mode="max", **kwargs)
[docs]
def avg_pool3d(inputs: SparseTensor, **kwargs) -> SparseTensor:
return pool3d(inputs, mode="avg", **kwargs)
[docs]
def pool_transpose3d(
inputs: SparseTensor,
target: SparseTensor | None = None,
*,
kernel_size=2,
stride=2,
padding=0,
dilation=1,
) -> SparseTensor:
"""Average coarse rows onto generated or explicit fine support."""
size = make_ntuple(kernel_size, ndim=3)
step = make_ntuple(stride, ndim=3)
pad = make_ntuple(padding, ndim=3)
spacing = make_ntuple(dilation, ndim=3)
if any(inputs.stride[index] % step[index] for index in range(3)):
raise ValueError("transpose stride must divide the input sparse stride")
output_stride = tuple(inputs.stride[index] // step[index] for index in range(3))
if target is None:
target_coords = build_transposed_output_coords(
inputs.coords,
kernel_size=size,
stride=step,
padding=pad,
dilation=spacing,
)
target = inputs.with_coordinates(
feats=inputs.feats.new_empty(
(target_coords.shape[0], inputs.feats.shape[1])
),
coords=target_coords,
stride=output_stride,
)
elif target.stride != output_stride:
raise ValueError(
f"target stride {target.stride} does not match transposed output "
f"stride {output_stride}"
)
relation = build_target_transposed_out_in_map(
inputs.coords,
target.coords,
kernel_size=size,
stride=step,
padding=pad,
dilation=spacing,
)
return target.replace(feats=_pool_features(inputs.feats, relation, "avg"))
[docs]
def trilinear_upsample3d(
inputs: SparseTensor,
target: SparseTensor | None = None,
*,
stride=2,
) -> SparseTensor:
"""Upsample sparse features with normalized trilinear interpolation."""
step = make_ntuple(stride, ndim=3)
if any(inputs.stride[index] % step[index] for index in range(3)):
raise ValueError("upsample stride must divide the input sparse stride")
output_stride = tuple(inputs.stride[index] // step[index] for index in range(3))
size = tuple(2 * value - 1 for value in step)
pad = tuple(value - 1 for value in step)
if target is None:
target_coords = build_transposed_output_coords(
inputs.coords,
kernel_size=size,
stride=step,
padding=pad,
)
target = inputs.with_coordinates(
feats=inputs.feats.new_empty(
(target_coords.shape[0], inputs.feats.shape[1])
),
coords=target_coords,
stride=output_stride,
)
elif target.stride != output_stride:
raise ValueError(
f"target stride {target.stride} does not match upsample output "
f"stride {output_stride}"
)
relation = build_target_transposed_out_in_map(
inputs.coords,
target.coords,
kernel_size=size,
stride=step,
padding=pad,
)
weights = _trilinear_weights(step, inputs.feats)
return target.replace(
feats=_weighted_pool_features(inputs.feats, relation, weights)
)
[docs]
def global_pool(
inputs: SparseTensor,
*,
mode: Literal["sum", "avg", "max"] = "sum",
batch_size: int | None = None,
) -> torch.Tensor:
"""Reduce sparse features independently for every declared batch."""
if mode not in {"sum", "avg", "max"}:
raise ValueError("global pool mode must be 'sum', 'avg', or 'max'")
batch_size = _batch_size(inputs, batch_size)
channels = int(inputs.feats.shape[1])
batch = inputs.coords[:, 0].to(torch.long)
if torch.any(batch < 0) or torch.any(batch >= batch_size):
raise ValueError("coordinate batch indices exceed the declared batch size")
counts = torch.bincount(batch, minlength=batch_size)
if mode == "max" and torch.any(counts == 0):
raise ValueError("global max pooling does not accept empty batches")
if mode in {"sum", "avg"}:
output = inputs.feats.new_zeros((batch_size, channels))
output.index_add_(0, batch, inputs.feats)
if mode == "avg":
output = output / counts.clamp_min(1).to(inputs.feats.dtype).unsqueeze(1)
return output
output = inputs.feats.new_full((batch_size, channels), -torch.inf)
rows = batch.view(-1, 1).expand(-1, channels)
return output.scatter_reduce_(
0, rows, inputs.feats, reduce="amax", include_self=True
)
[docs]
def global_sum_pool(
inputs: SparseTensor, *, batch_size: int | None = None
) -> torch.Tensor:
return global_pool(inputs, mode="sum", batch_size=batch_size)
[docs]
def global_avg_pool(
inputs: SparseTensor, *, batch_size: int | None = None
) -> torch.Tensor:
return global_pool(inputs, mode="avg", batch_size=batch_size)
[docs]
def global_max_pool(
inputs: SparseTensor, *, batch_size: int | None = None
) -> torch.Tensor:
return global_pool(inputs, mode="max", batch_size=batch_size)
def _batch_size(inputs: SparseTensor, explicit: int | None) -> int:
inferred = (
len(inputs.batch_counts)
if inputs.batch_counts is not None
else int(inputs.spatial_range[0])
if inputs.spatial_range is not None
else int(inputs.coords[:, 0].max().item()) + 1
if inputs.coords.shape[0] > 0
else 0
)
if explicit is None:
return inferred
if explicit < inferred:
raise ValueError("batch_size is smaller than the sparse batch metadata")
return int(explicit)
def _pool_features(
feats: torch.Tensor, relation: torch.Tensor, mode: PoolMode
) -> torch.Tensor:
output_size = int(relation.shape[0])
channels = int(feats.shape[1])
valid = relation >= 0
if not torch.any(valid):
return feats.new_zeros((output_size, channels))
out_rows = torch.nonzero(valid, as_tuple=False)[:, 0]
in_rows = relation[valid].to(torch.long)
gathered = feats.index_select(0, in_rows)
if mode in {"sum", "avg"}:
output = feats.new_zeros((output_size, channels))
output.index_add_(0, out_rows, gathered)
if mode == "avg":
counts = valid.sum(dim=1).clamp_min(1).to(feats.dtype).unsqueeze(1)
output = output / counts
return output
output = feats.new_full((output_size, channels), -torch.inf)
scatter_rows = out_rows.view(-1, 1).expand(-1, channels)
output.scatter_reduce_(0, scatter_rows, gathered, reduce="amax", include_self=True)
empty = ~torch.any(valid, dim=1)
if torch.any(empty):
output[empty] = 0
return output
def _weighted_pool_features(
feats: torch.Tensor,
relation: torch.Tensor,
weights: torch.Tensor,
) -> torch.Tensor:
valid = relation >= 0
gathered = feats.index_select(
0, relation.clamp_min(0).reshape(-1).to(torch.long)
).reshape(*relation.shape, feats.shape[1])
coefficients = valid.to(feats.dtype) * weights.unsqueeze(0)
summed = torch.sum(gathered * coefficients.unsqueeze(2), dim=1)
denominator = torch.sum(coefficients, dim=1, keepdim=True)
return torch.where(
denominator > 0,
summed / denominator.clamp_min(torch.finfo(feats.dtype).tiny),
summed,
)
def _trilinear_weights(
stride: tuple[int, int, int], reference: torch.Tensor
) -> torch.Tensor:
axes = tuple(
tuple(1.0 - abs(offset - (step - 1)) / step for offset in range(2 * step - 1))
for step in stride
)
return reference.new_tensor([x * y * z for x, y, z in product(*axes)])
def _pooled_spatial_range(spatial_range, kernel_size, stride, padding, dilation):
if spatial_range is None:
return None
return tuple(spatial_range[:1]) + tuple(
max(
0,
(
int(spatial_range[index + 1])
+ 2 * padding[index]
- dilation[index] * (kernel_size[index] - 1)
- 1
)
// stride[index]
+ 1,
)
for index in range(3)
)