torch_harmonics.NeighborhoodAttentionS2#

class torch_harmonics.NeighborhoodAttentionS2(
in_channels,
in_shape,
out_shape,
grid_in='equiangular',
grid_out='equiangular',
num_heads=1,
scale=None,
use_qknorm=False,
bias=True,
theta_cutoff=None,
k_channels=None,
out_channels=None,
optimized_kernel=True,
)[source]#

Bases: Module

Neighborhood attention on the 2-sphere.

This is the local counterpart of AttentionS2. Instead of attending globally, every output location attends only to the input points inside a geodesic neighborhood around it – the spherical disk \(D(x) = \{x' \in S^2 : d(x, x') \le \theta_\mathrm{cutoff}\}\), where \(d(\cdot, \cdot)\) is the great-circle (Haversine) distance and \(\theta_\mathrm{cutoff}\) the cutoff radius. Restricting attention to this disk adds an inductive bias for locality and lowers the cost from \(\mathcal{O}(N^2)\) to \(\mathcal{O}(k N)\), where \(k\) is the number of points in a neighborhood.

Following [4], the attention softmax integrates over the neighborhood against the sphere’s numerical quadrature weights. This makes the layer a resolution-agnostic neural operator – it can be evaluated on arbitrary grid resolutions (though the learned features themselves remain resolution dependent) – and approximately \(SO(3)\)-equivariant, since the underlying integrals are invariant under rotations (the Haar measure).

The sparse neighborhood structure is precomputed with the same discrete-continuous construction used for the DISCO convolutions (DiscreteContinuousConvS2): Here, only the suppot (index information) of the zero order DISCO kernel is used to define an indicator function of the cutoff disk, so that any input point contributes to an output location exactly when it lies within \(\theta_\mathrm{cutoff}\) of it. The relative weight of each input point depends on their contribution to the softmax as well as their quadrature weights.

Parameters:
  • in_channels (int) – number of channels of the input signal (corresponds to embed_dim in MHA in PyTorch)

  • in_shape (tuple) – shape of the input grid

  • out_shape (tuple) – shape of the output grid

  • grid_in (str, optional) – input grid type, "equiangular" by default

  • grid_out (str, optional) – output grid type, "equiangular" by default

  • bias (bool, optional) – if specified, adds bias to input / output projection layers

  • theta_cutoff (float, optional) – Angular radius of the geodesic neighborhood disk, in radians. Input points farther than this from an output location are excluded from its attention. If None (default), it is set to one latitudinal grid spacing of the coarser of the input and output grids, i.e. pi / (nlat - 1). Must be positive.

  • k_channels (int) – number of dimensions for interior inner product in the attention matrix (corresponds to kdim in MHA in PyTorch)

  • out_channels (int, optional) – number of dimensions for interior inner product in the attention matrix (corresponds to vdim in MHA in PyTorch)

  • optimized_kernel (Optional[bool]) – Whether to use the optimized kernel (if available)

  • num_heads (int | None)

  • scale (Tensor | float | None)

  • use_qknorm (bool | None)

References

[4]

forward(query, key=None, value=None)[source]#

Apply neighborhood attention on the sphere.

Parameters:
  • query (torch.Tensor) – Query signal of shape (batch, in_channels, nlat_out, nlon_out) (sampled on the output grid).

  • key (torch.Tensor, optional) – Key signal of shape (batch, in_channels, nlat_in, nlon_in) (sampled on the input grid). Defaults to query (self-attention, which requires matching input and output grids).

  • value (torch.Tensor, optional) – Value signal of shape (batch, in_channels, nlat_in, nlon_in) (sampled on the input grid). Defaults to query (self-attention, which requires matching input and output grids).

Returns:

Attention output of shape (batch, out_channels, nlat_out, nlon_out).

Return type:

torch.Tensor