torch_harmonics.AttentionS2#

class torch_harmonics.AttentionS2(
in_channels,
num_heads,
in_shape,
out_shape,
grid_in='equiangular',
grid_out='equiangular',
scale=None,
use_qknorm=False,
bias=True,
k_channels=None,
out_channels=None,
drop_rate=0.0,
)[source]#

Bases: Module

(Global) attention on the 2-sphere.

This is ordinary (global) scaled dot-product attention, made geometrically faithful on the sphere by folding the numerical quadrature weights of the grid into the attention. Following [4], the softmax over keys becomes a quadrature approximation of a continuous attention integral over the sphere: the logarithms of the spherical quadrature weights are added to the pre-softmax attention scores as an additive mask, so that after the softmax exponential they act as multiplicative quadrature weights in the normalization. Using log-weights lets them be passed directly as the attn_mask of torch.nn.functional.scaled_dot_product_attention().

Incorporating the quadrature weights this way makes the layer a resolution-agnostic neural operator (evaluable on arbitrary grids, though the learned features remain resolution dependent) and approximately \(SO(3)\)-equivariant, since the underlying integral is invariant under rotations (the Haar measure). For the local variant that confines attention to a geodesic neighborhood, see NeighborhoodAttentionS2.

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

  • num_heads (int) – number of attention heads

  • 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

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

  • scale (Tensor | float | None)

  • use_qknorm (bool | None)

  • drop_rate (float | None)

References

[4]

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

Apply global 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). Defaults to query (self-attention).

  • value (torch.Tensor, optional) – Value signal of shape (batch, in_channels, nlat_in, nlon_in). Defaults to query (self-attention).

Returns:

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

Return type:

torch.Tensor