Source code for torch_lattice.nn.functional.pooling

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) )