torch_harmonics.distributed.DistributedNeighborhoodAttentionS2#

class torch_harmonics.distributed.DistributedNeighborhoodAttentionS2(
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: NeighborhoodAttentionS2

Distributed neighborhood attention on the 2-sphere using a ring exchange strategy for the longitude dimension and halo exchange for the latitude dimension.

Data is assumed to be split along both the latitude (polar group) and longitude (azimuth group) dimensions. The forward pass uses ring exchange of key/value chunks over the azimuth group so that every output point can attend to its full spherical neighborhood.

All three directions of the serial layer are supported: self-attention (in_shape == out_shape), downsampling cross-attention (gather kernels, nlon_in % nlon_out == 0) and upsampling cross-attention (scatter kernels, nlon_out % nlon_in == 0). In all cases K/V (which live on the input grid) rotate around the azimuth ring while Q and the softmax state stay local.

Inherits learnable parameters from torch_harmonics.NeighborhoodAttentionS2.

See also

torch_harmonics.NeighborhoodAttentionS2

Serial counterpart with full parameter documentation.

Parameters: