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,
Bases:
NeighborhoodAttentionS2Distributed 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.NeighborhoodAttentionS2Serial counterpart with full parameter documentation.