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,
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_maskoftorch.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 defaultgrid_out (str, optional) – output grid type,
"equiangular"by defaultbias (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)
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 toquery(self-attention).value (torch.Tensor, optional) – Value signal of shape
(batch, in_channels, nlat_in, nlon_in). Defaults toquery(self-attention).
- Returns:
Attention output of shape
(batch, out_channels, nlat_out, nlon_out).- Return type: