# coding=utf-8
# SPDX-FileCopyrightText: Copyright (c) 2022 The torch-harmonics Authors. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
#
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
#
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
#
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
#
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
#
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
from itertools import accumulate
from typing import Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.nn as nn
from attention_helpers import optimized_kernels_is_available
from torch_harmonics.attention import attention_kernels
from torch_harmonics.attention.attention import NeighborhoodAttentionS2
from torch_harmonics.distributed._amp_utils import _cast_to_autocast_dtype, _custom_setup_context
from .primitives import compute_split_shapes, get_group_neighbors, polar_halo_exchange
from .utils import azimuth_group, azimuth_group_rank, azimuth_group_size, polar_group_rank, polar_group_size
# ---------------------------------------------------------------------------
# autograd.Function for the ring-step attention kernel calls
# ---------------------------------------------------------------------------
@torch.compiler.disable()
def _ring_kv(kw_chunk, vw_chunk, az_group, next_nlon_kw, next_nlon_kv):
"""Async send current chunks, receive next chunks with known shapes."""
send_to, recv_from = get_group_neighbors(az_group)
B, C_k, H, _ = kw_chunk.shape
B, C_v, H, _ = vw_chunk.shape
recv_kw = torch.empty(B, C_k, H, next_nlon_kw, device=kw_chunk.device, dtype=kw_chunk.dtype)
recv_vw = torch.empty(B, C_v, H, next_nlon_kv, device=vw_chunk.device, dtype=vw_chunk.dtype)
ops = [
dist.P2POp(dist.isend, kw_chunk, send_to, az_group),
dist.P2POp(dist.irecv, recv_kw, recv_from, az_group),
dist.P2POp(dist.isend, vw_chunk, send_to, az_group),
dist.P2POp(dist.irecv, recv_vw, recv_from, az_group),
]
reqs = dist.batch_isend_irecv(ops)
return recv_kw, recv_vw, reqs
class _RingNeighborhoodAttentionFn(torch.autograd.Function):
"""Forward ring attention + backward ring for one attention head group.
kw, vw : [B*nh, C_k/C_v, H_halo, W_local] channels-first, lat-halo-padded
qw : [B*nh, C_k, H_out_local, W_out_local] channels-first
State buffers use channels-last layout as required by the CUDA kernels:
y_acc : [B, H_out, W_out, C_v]
alpha_k/kvw : [B, H_out, W_out, C_k]
alpha_sum/qdotk_max/integral : [B, H_out, W_out]
"""
@staticmethod
@torch.amp.custom_fwd(device_type="cuda")
def forward(
kw,
vw,
qw,
psi_col_idx,
psi_roff_idx,
psi_row_idx,
quad_weights,
nlon_in: int,
pscale: int,
lon_chunk_starts: list,
nlon_kx_list: list,
lat_halo_start: int,
nlat_out_local: int,
nlon_out_local: int,
r_lat: int,
az_group,
az_rank: int,
az_size: int,
psi_n_long_rows: int,
psi_max_row_len: int,
psi_mid_row_len: int,
):
B, _, _, _ = kw.shape
_, C_v, _, _ = vw.shape
device = kw.device
# Capture input dtype so we can cast the user-visible output (y_out)
# back at the end of forward. The internal accumulators (y_acc,
# alpha_sum, qdotk_max) stay fp32 — softmax stability requires that,
# and the saved alpha_sum/qdotk_max feed fp32 math in backward.
inp_dtype = kw.dtype
# Allocate state buffers in formats expected by the CUDA kernels:
# y_acc: channels-last [B, H, W, C_v]; scalars: [B, H, W]
y_acc = torch.zeros(B, nlat_out_local, nlon_out_local, C_v, device=device, dtype=torch.float32)
alpha_sum = torch.zeros(B, nlat_out_local, nlon_out_local, device=device, dtype=torch.float32)
qdotk_max = torch.full((B, nlat_out_local, nlon_out_local), float("-inf"), device=device, dtype=torch.float32)
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
# Pre-allocate receive buffers for the NEXT chunk (correct shape)
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
attention_kernels.forward_ring_step.default(
kw_chunk,
vw_chunk,
qw,
y_acc,
alpha_sum,
qdotk_max,
quad_weights,
psi_col_idx,
psi_roff_idx,
psi_row_idx,
nlon_in,
pscale,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
psi_n_long_rows,
psi_max_row_len,
psi_mid_row_len,
)
if step < az_size - 1:
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# Finalize: y = y_acc / alpha_sum (both channels-last layout). Cast
# back to the input dtype to keep the op faithful to its input dtype.
y_out = y_acc / alpha_sum.unsqueeze(-1) # [B, H, W, C_v]
y_out = y_out.permute(0, 3, 1, 2).to(dtype=inp_dtype).contiguous() # [B, C_v, H, W]
# alpha_sum and qdotk_max are returned so setup_context can save them;
# they are marked non-differentiable there, so backward still only
# receives one gradient argument (dy for y_out).
return y_out, alpha_sum, qdotk_max
@staticmethod
@_custom_setup_context(device_type="cuda")
def setup_context(ctx, inputs, output):
(
kw,
vw,
qw,
psi_col_idx,
psi_roff_idx,
psi_row_idx,
quad_weights,
nlon_in,
pscale,
lon_chunk_starts,
nlon_kx_list,
lat_halo_start,
nlat_out_local,
nlon_out_local,
r_lat,
az_group,
az_rank,
az_size,
psi_n_long_rows,
psi_max_row_len,
psi_mid_row_len,
) = inputs
y_out, alpha_sum, qdotk_max = output
# alpha_sum and qdotk_max are internal accumulators, not true outputs;
# marking them non-differentiable keeps backward's signature as (ctx, dy).
ctx.mark_non_differentiable(alpha_sum, qdotk_max)
ctx.save_for_backward(kw, vw, qw, psi_col_idx, psi_roff_idx, psi_row_idx, quad_weights, alpha_sum, qdotk_max)
ctx.nlon_in = nlon_in
ctx.pscale = pscale
ctx.lon_chunk_starts = lon_chunk_starts
ctx.nlon_kx_list = nlon_kx_list
ctx.lat_halo_start = lat_halo_start
ctx.nlat_out_local = nlat_out_local
ctx.nlon_out_local = nlon_out_local
ctx.az_group = az_group
ctx.az_rank = az_rank
ctx.az_size = az_size
ctx.psi_n_long_rows = psi_n_long_rows
ctx.psi_max_row_len = psi_max_row_len
ctx.psi_mid_row_len = psi_mid_row_len
@staticmethod
@torch.amp.custom_bwd(device_type="cuda")
def backward(ctx, dy, _dalpha_sum, _dqdotk_max):
# _dalpha_sum and _dqdotk_max are always None (non-differentiable outputs)
(kw, vw, qw, psi_col_idx, psi_roff_idx, psi_row_idx, quad_weights, fwd_alpha_sum, fwd_qdotk_max) = ctx.saved_tensors
nlon_in = ctx.nlon_in
pscale = ctx.pscale
lon_chunk_starts = ctx.lon_chunk_starts
nlon_kx_list = ctx.nlon_kx_list
lat_halo_start = ctx.lat_halo_start
nlat_out_local = ctx.nlat_out_local
nlon_out_local = ctx.nlon_out_local
az_group = ctx.az_group
az_rank = ctx.az_rank
az_size = ctx.az_size
psi_n_long_rows = ctx.psi_n_long_rows
psi_max_row_len = ctx.psi_max_row_len
psi_mid_row_len = ctx.psi_mid_row_len
# Autograd contract: skip per-branch work (kernel calls, allreduces) for any
# of {kw, vw, qw} that doesn't need a gradient, and return None in those slots.
# This is what lets torch.compile / AOTAutograd prune dead subgraphs (including
# the NCCL allreduces) from the compiled backward.
kw_needs_grad = ctx.needs_input_grad[0]
vw_needs_grad = ctx.needs_input_grad[1]
qw_needs_grad = ctx.needs_input_grad[2]
# Defensive: if somehow none of (kw, vw, qw) need grad (e.g., user wired
# requires_grad onto one of the index buffers), there's nothing to compute.
if not (kw_needs_grad or vw_needs_grad or qw_needs_grad):
return (None,) * 21
B, C_k, H_halo, _ = kw.shape
_, C_v, _, _ = vw.shape
device = kw.device
# Capture input dtypes so the returned grads can be cast back. The
# backward ring kernels consume kw/vw/qw/dy in their native dtype (Tier B:
# widen fp16/bf16 at load, fp32 compute/accumulation). Keeping them native
# — instead of upcasting to fp32 here — also keeps the backward ring
# exchange at 16-bit under AMP (halved K/V comm volume), matching the
# forward ring. The fp32 accumulators (integral_buf, alpha_k/kvw_buf,
# dkw/dvw_full_cl) are unaffected; the returned grads are cast back to the
# captured input dtypes at the end.
kw_dtype = kw.dtype
vw_dtype = vw.dtype
qw_dtype = qw.dtype
dy_cf = dy.contiguous() # channels-first [B, C_v, H, W], native dtype
# ----------------------------------------------------------------
# Backward pass 1: re-accumulate {alpha_sum, qdotk_max, integral,
# alpha_k, alpha_kvw} via ring.
# Required whenever any of (kw, vw, qw) needs grad: integral feeds
# pass-2's integral_norm and dqy reads alpha_k/alpha_kvw. The kernel
# writes all three buffers in one call, so pass-1 cannot be pruned
# per-branch.
# ----------------------------------------------------------------
bwd_alpha_sum = torch.zeros(B, nlat_out_local, nlon_out_local, device=device, dtype=torch.float32)
bwd_qdotk_max = torch.full((B, nlat_out_local, nlon_out_local), float("-inf"), device=device, dtype=torch.float32)
integral_buf = torch.zeros_like(bwd_alpha_sum)
alpha_k_buf = torch.zeros(B, nlat_out_local, nlon_out_local, C_k, device=device, dtype=torch.float32)
alpha_kvw_buf = torch.zeros_like(alpha_k_buf)
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
attention_kernels.backward_ring_step_pass1.default(
kw_chunk,
vw_chunk,
qw,
dy_cf,
bwd_alpha_sum,
bwd_qdotk_max,
integral_buf,
alpha_k_buf,
alpha_kvw_buf,
quad_weights,
psi_col_idx,
psi_roff_idx,
psi_row_idx,
nlon_in,
pscale,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
psi_n_long_rows,
psi_max_row_len,
psi_mid_row_len,
)
if step < az_size - 1:
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# ----------------------------------------------------------------
# Finalize pass-1 outputs.
# Use the SAVED forward alpha_sum/qdotk_max (same values, but authoritative).
# ----------------------------------------------------------------
alpha_sum_inv = 1.0 / fwd_alpha_sum # [B, H, W]
# integral_norm only feeds pass-2; skip if neither kw nor vw needs grad.
if kw_needs_grad or vw_needs_grad:
integral_norm = integral_buf * alpha_sum_inv # [B, H, W]
# dqy[b,h,w,c] = inv_sq*(alpha_sum*alpha_kvw - integral*alpha_k)
if qw_needs_grad:
alpha_sum_inv_sq = alpha_sum_inv**2
dqy_cl = alpha_sum_inv_sq.unsqueeze(-1) * (fwd_alpha_sum.unsqueeze(-1) * alpha_kvw_buf - integral_buf.unsqueeze(-1) * alpha_k_buf) # [B, H, W, C_k]
dqy = dqy_cl.permute(0, 3, 1, 2).to(dtype=qw_dtype).contiguous() # [B, C_k, H, W]
else:
dqy = None
# ----------------------------------------------------------------
# Backward pass 2: scatter dkw/dvw contributions.
# Each GPU computes its contribution to every lon chunk it visits;
# then allreduce across azimuth ranks, extract local chunk.
# Skip entirely if neither kw nor vw needs grad. The fused kernel
# writes both dkw_chunk_cl and dvw_chunk_cl in one call, so the
# per-chunk allocations stay; we just gate the accumulation /
# allreduce / extract per branch.
# TODO: replace allreduce with ring reduce-scatter for efficiency.
# ----------------------------------------------------------------
if kw_needs_grad or vw_needs_grad:
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
nlon_in_total = sum(nlon_kx_list)
dkw_full_cl = torch.zeros(B, H_halo, nlon_in_total, C_k, device=device, dtype=torch.float32) if kw_needs_grad else None
dvw_full_cl = torch.zeros(B, H_halo, nlon_in_total, C_v, device=device, dtype=torch.float32) if vw_needs_grad else None
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
nlon_kx = nlon_kx_list[src_rank]
# Channels-last gradient buffers for this chunk (both required by the
# fused kernel signature; we discard the one we don't need).
dkw_chunk_cl = torch.zeros(B, H_halo, nlon_kx, C_k, device=device, dtype=torch.float32)
dvw_chunk_cl = torch.zeros(B, H_halo, nlon_kx, C_v, device=device, dtype=torch.float32)
attention_kernels.backward_ring_step_pass2.default(
kw_chunk,
vw_chunk,
qw,
dy_cf,
fwd_alpha_sum,
fwd_qdotk_max,
integral_norm,
dkw_chunk_cl,
dvw_chunk_cl,
quad_weights,
psi_col_idx,
psi_roff_idx,
psi_row_idx,
nlon_in,
pscale,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
psi_n_long_rows,
psi_max_row_len,
psi_mid_row_len,
)
if kw_needs_grad:
dkw_full_cl[:, :, lon_lo_kx : lon_lo_kx + nlon_kx, :].add_(dkw_chunk_cl)
if vw_needs_grad:
dvw_full_cl[:, :, lon_lo_kx : lon_lo_kx + nlon_kx, :].add_(dvw_chunk_cl)
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# Per-branch allreduce — only the branches we'll return.
if az_size > 1 and az_group is not None:
if kw_needs_grad:
dist.all_reduce(dkw_full_cl, group=az_group)
if vw_needs_grad:
dist.all_reduce(dvw_full_cl, group=az_group)
my_lo = lon_chunk_starts[az_rank]
my_nlon = nlon_kx_list[az_rank]
# Extract local chunk and convert channels-last → channels-first.
# No halo stripping: dkw/dvw must match kw/vw shape (= key_halo/value_halo).
# The autograd through torch.cat in _exchange_lat_halo extracts the
# middle H_in rows as the gradient for key_proj/value_proj.
if kw_needs_grad:
dkw_cl = dkw_full_cl[:, :, my_lo : my_lo + my_nlon, :].contiguous()
dkw = dkw_cl.permute(0, 3, 1, 2).to(dtype=kw_dtype).contiguous() # [B, C_k, H_halo, W_local]
else:
dkw = None
if vw_needs_grad:
dvw_cl = dvw_full_cl[:, :, my_lo : my_lo + my_nlon, :].contiguous()
dvw = dvw_cl.permute(0, 3, 1, 2).to(dtype=vw_dtype).contiguous() # [B, C_v, H_halo, W_local]
else:
dvw = None
else:
dkw = None
dvw = None
# Return grads for (kw, vw, qw, psi_col, psi_roff, psi_row, quad_weights,
# nlon_in, pscale, lon_chunk_starts, nlon_kx_list, lat_halo_start,
# nlat_out_local, nlon_out_local, r_lat,
# az_group, az_rank, az_size,
# psi_n_long_rows, psi_max_row_len, psi_mid_row_len)
return (dkw, dvw, dqy, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None)
class _RingNeighborhoodAttentionUpsampleFn(torch.autograd.Function):
"""Forward ring attention + backward ring for the UPSAMPLE (scatter) direction.
K/V live on the coarse input grid (sharded, halo-padded in lat, rotating
around the azimuth ring); Q and the output live on the fine output grid and
stay local.
kw, vw : [B*nh, C_k/C_v, H_halo, W_in_local] channels-first, lat-halo-padded
qw : [B*nh, C_k, H_out_local, W_out_local] channels-first
State buffers use channels-last layout as required by the CUDA kernels:
y_acc : [B, H_out, W_out, C_v]
alpha_k/kvw : [B, H_out, W_out, C_k]
alpha_sum/qdotk_max/integral : [B, H_out, W_out]
The local psi is built by _build_local_psi_upsample: rows are keyed by the
halo-padded LOCAL input latitude, cols encode (ho_local, wo_shifted) on the
fine output grid with wo pre-shifted by -lon_lo_out.
"""
@staticmethod
@torch.amp.custom_fwd(device_type="cuda")
def forward(
kw,
vw,
qw,
psi_col_idx,
psi_roff_idx,
quad_weights,
nlon_in: int,
nlon_out_global: int,
pscale_out: int,
lon_chunk_starts: list,
nlon_kx_list: list,
lat_halo_start: int,
nlat_out_local: int,
nlon_out_local: int,
r_lat: int,
az_group,
az_rank: int,
az_size: int,
):
B, _, _, _ = kw.shape
_, C_v, _, _ = vw.shape
device = kw.device
# Capture input dtype so we can cast the user-visible output (y_out)
# back at the end of forward. The internal accumulators stay fp32 —
# softmax stability requires that, and the saved alpha_sum/qdotk_max
# feed fp32 math in backward.
inp_dtype = kw.dtype
# Allocate state buffers in formats expected by the CUDA kernels:
# y_acc: channels-last [B, H, W, C_v]; scalars: [B, H, W]
y_acc = torch.zeros(B, nlat_out_local, nlon_out_local, C_v, device=device, dtype=torch.float32)
alpha_sum = torch.zeros(B, nlat_out_local, nlon_out_local, device=device, dtype=torch.float32)
qdotk_max = torch.full((B, nlat_out_local, nlon_out_local), float("-inf"), device=device, dtype=torch.float32)
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
# Pre-allocate receive buffers for the NEXT chunk (correct shape)
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
attention_kernels.forward_ring_step_upsample.default(
kw_chunk,
vw_chunk,
qw,
y_acc,
alpha_sum,
qdotk_max,
quad_weights,
psi_col_idx,
psi_roff_idx,
nlon_in,
nlon_out_global,
pscale_out,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
)
if step < az_size - 1:
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# Finalize: y = y_acc / alpha_sum (both channels-last layout). Cast
# back to the input dtype to keep the op faithful to its input dtype.
y_out = y_acc / alpha_sum.unsqueeze(-1) # [B, H, W, C_v]
y_out = y_out.permute(0, 3, 1, 2).to(dtype=inp_dtype).contiguous() # [B, C_v, H, W]
# alpha_sum and qdotk_max are returned so setup_context can save them;
# they are marked non-differentiable there, so backward still only
# receives one gradient argument (dy for y_out).
return y_out, alpha_sum, qdotk_max
@staticmethod
@_custom_setup_context(device_type="cuda")
def setup_context(ctx, inputs, output):
(
kw,
vw,
qw,
psi_col_idx,
psi_roff_idx,
quad_weights,
nlon_in,
nlon_out_global,
pscale_out,
lon_chunk_starts,
nlon_kx_list,
lat_halo_start,
nlat_out_local,
nlon_out_local,
r_lat,
az_group,
az_rank,
az_size,
) = inputs
y_out, alpha_sum, qdotk_max = output
# alpha_sum and qdotk_max are internal accumulators, not true outputs;
# marking them non-differentiable keeps backward's signature as (ctx, dy).
ctx.mark_non_differentiable(alpha_sum, qdotk_max)
ctx.save_for_backward(kw, vw, qw, psi_col_idx, psi_roff_idx, quad_weights, alpha_sum, qdotk_max)
ctx.nlon_in = nlon_in
ctx.nlon_out_global = nlon_out_global
ctx.pscale_out = pscale_out
ctx.lon_chunk_starts = lon_chunk_starts
ctx.nlon_kx_list = nlon_kx_list
ctx.lat_halo_start = lat_halo_start
ctx.nlat_out_local = nlat_out_local
ctx.nlon_out_local = nlon_out_local
ctx.az_group = az_group
ctx.az_rank = az_rank
ctx.az_size = az_size
@staticmethod
@torch.amp.custom_bwd(device_type="cuda")
def backward(ctx, dy, _dalpha_sum, _dqdotk_max):
# _dalpha_sum and _dqdotk_max are always None (non-differentiable outputs)
(kw, vw, qw, psi_col_idx, psi_roff_idx, quad_weights, fwd_alpha_sum, fwd_qdotk_max) = ctx.saved_tensors
nlon_in = ctx.nlon_in
nlon_out_global = ctx.nlon_out_global
pscale_out = ctx.pscale_out
lon_chunk_starts = ctx.lon_chunk_starts
nlon_kx_list = ctx.nlon_kx_list
lat_halo_start = ctx.lat_halo_start
nlat_out_local = ctx.nlat_out_local
nlon_out_local = ctx.nlon_out_local
az_group = ctx.az_group
az_rank = ctx.az_rank
az_size = ctx.az_size
# Autograd contract: skip per-branch work (kernel calls, allreduces) for any
# of {kw, vw, qw} that doesn't need a gradient, and return None in those slots.
kw_needs_grad = ctx.needs_input_grad[0]
vw_needs_grad = ctx.needs_input_grad[1]
qw_needs_grad = ctx.needs_input_grad[2]
# Defensive: if somehow none of (kw, vw, qw) need grad, there's nothing to compute.
if not (kw_needs_grad or vw_needs_grad or qw_needs_grad):
return (None,) * 18
B, C_k, H_halo, _ = kw.shape
_, C_v, _, _ = vw.shape
device = kw.device
# Capture input dtypes so the returned grads can be cast back. The
# backward kernels consume kw/vw/qw/dy in their native dtype (widen at
# load, fp32 compute/accumulation), keeping the backward ring exchange
# at 16-bit under AMP.
kw_dtype = kw.dtype
vw_dtype = vw.dtype
qw_dtype = qw.dtype
dy_cf = dy.contiguous() # channels-first [B, C_v, H, W], native dtype
# ----------------------------------------------------------------
# Backward pass 1: re-accumulate {integral, alpha_k, alpha_kvw} via
# ring, using the SAVED forward alpha_sum/qdotk_max (no max recompute
# is needed in the upsample direction — the forward-final softmax
# stats are authoritative). Required whenever any of (kw, vw, qw)
# needs grad: integral feeds pass-2's integral_norm and dqy reads
# alpha_k/alpha_kvw. The kernel writes all buffers in one call, so
# pass-1 cannot be pruned per-branch.
# ----------------------------------------------------------------
integral_buf = torch.zeros(B, nlat_out_local, nlon_out_local, device=device, dtype=torch.float32)
alpha_k_buf = torch.zeros(B, nlat_out_local, nlon_out_local, C_k, device=device, dtype=torch.float32)
alpha_kvw_buf = torch.zeros_like(alpha_k_buf)
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
attention_kernels.backward_ring_step_upsample_pass1.default(
kw_chunk,
vw_chunk,
qw,
dy_cf,
fwd_qdotk_max,
integral_buf,
alpha_k_buf,
alpha_kvw_buf,
quad_weights,
psi_col_idx,
psi_roff_idx,
nlon_in,
nlon_out_global,
pscale_out,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
)
if step < az_size - 1:
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# ----------------------------------------------------------------
# Finalize pass-1 outputs.
# ----------------------------------------------------------------
alpha_sum_inv = 1.0 / fwd_alpha_sum # [B, H, W]
# integral_norm only feeds pass-2; skip if neither kw nor vw needs grad.
if kw_needs_grad or vw_needs_grad:
integral_norm = integral_buf * alpha_sum_inv # [B, H, W]
# dqy[b,h,w,c] = inv_sq*(alpha_sum*alpha_kvw - integral*alpha_k)
if qw_needs_grad:
alpha_sum_inv_sq = alpha_sum_inv**2
dqy_cl = alpha_sum_inv_sq.unsqueeze(-1) * (fwd_alpha_sum.unsqueeze(-1) * alpha_kvw_buf - integral_buf.unsqueeze(-1) * alpha_k_buf) # [B, H, W, C_k]
dqy = dqy_cl.permute(0, 3, 1, 2).to(dtype=qw_dtype).contiguous() # [B, C_k, H, W]
else:
dqy = None
# ----------------------------------------------------------------
# Backward pass 2: accumulate dkw/dvw contributions per chunk.
# Each GPU computes its LOCAL outputs' contribution to every lon chunk
# it visits; then allreduce across azimuth ranks, extract local chunk.
# TODO: replace allreduce with ring reduce-scatter for efficiency.
# ----------------------------------------------------------------
if kw_needs_grad or vw_needs_grad:
kw_chunk = kw.contiguous()
vw_chunk = vw.contiguous()
nlon_in_total = sum(nlon_kx_list)
dkw_full_cl = torch.zeros(B, H_halo, nlon_in_total, C_k, device=device, dtype=torch.float32) if kw_needs_grad else None
dvw_full_cl = torch.zeros(B, H_halo, nlon_in_total, C_v, device=device, dtype=torch.float32) if vw_needs_grad else None
for step in range(az_size):
src_rank = (az_rank + step) % az_size
lon_lo_kx = lon_chunk_starts[src_rank]
nlon_kx = nlon_kx_list[src_rank]
# Channels-last gradient buffers for this chunk (both required by the
# fused kernel signature; we discard the one we don't need).
dkw_chunk_cl = torch.zeros(B, H_halo, nlon_kx, C_k, device=device, dtype=torch.float32)
dvw_chunk_cl = torch.zeros(B, H_halo, nlon_kx, C_v, device=device, dtype=torch.float32)
attention_kernels.backward_ring_step_upsample_pass2.default(
kw_chunk,
vw_chunk,
qw,
dy_cf,
fwd_alpha_sum,
fwd_qdotk_max,
integral_norm,
dkw_chunk_cl,
dvw_chunk_cl,
quad_weights,
psi_col_idx,
psi_roff_idx,
nlon_in,
nlon_out_global,
pscale_out,
lon_lo_kx,
lat_halo_start,
nlat_out_local,
nlon_out_local,
)
if kw_needs_grad:
dkw_full_cl[:, :, lon_lo_kx : lon_lo_kx + nlon_kx, :].add_(dkw_chunk_cl)
if vw_needs_grad:
dvw_full_cl[:, :, lon_lo_kx : lon_lo_kx + nlon_kx, :].add_(dvw_chunk_cl)
if step < az_size - 1:
next_src = (az_rank + step + 1) % az_size
recv_kw, recv_vw, reqs = _ring_kv(kw_chunk, vw_chunk, az_group, nlon_kx_list[next_src], nlon_kx_list[next_src])
for req in reqs:
req.wait()
kw_chunk = recv_kw.clone()
vw_chunk = recv_vw.clone()
# Per-branch allreduce — only the branches we'll return.
if az_size > 1 and az_group is not None:
if kw_needs_grad:
dist.all_reduce(dkw_full_cl, group=az_group)
if vw_needs_grad:
dist.all_reduce(dvw_full_cl, group=az_group)
my_lo = lon_chunk_starts[az_rank]
my_nlon = nlon_kx_list[az_rank]
# Extract local chunk and convert channels-last → channels-first.
# No halo stripping: dkw/dvw must match kw/vw shape (= key_halo/value_halo).
if kw_needs_grad:
dkw_cl = dkw_full_cl[:, :, my_lo : my_lo + my_nlon, :].contiguous()
dkw = dkw_cl.permute(0, 3, 1, 2).to(dtype=kw_dtype).contiguous() # [B, C_k, H_halo, W_local]
else:
dkw = None
if vw_needs_grad:
dvw_cl = dvw_full_cl[:, :, my_lo : my_lo + my_nlon, :].contiguous()
dvw = dvw_cl.permute(0, 3, 1, 2).to(dtype=vw_dtype).contiguous() # [B, C_v, H_halo, W_local]
else:
dvw = None
else:
dkw = None
dvw = None
# Return grads for (kw, vw, qw, psi_col, psi_roff, quad_weights,
# nlon_in, nlon_out_global, pscale_out, lon_chunk_starts,
# nlon_kx_list, lat_halo_start, nlat_out_local, nlon_out_local,
# r_lat, az_group, az_rank, az_size)
return (dkw, dvw, dqy, None, None, None, None, None, None, None, None, None, None, None, None, None, None, None)
# ---------------------------------------------------------------------------
# Distributed Neighborhood Attention on the 2-sphere
# ---------------------------------------------------------------------------
[docs]
class DistributedNeighborhoodAttentionS2(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 :class:`torch_harmonics.NeighborhoodAttentionS2`.
.. seealso::
:class:`torch_harmonics.NeighborhoodAttentionS2`
Serial counterpart with full parameter documentation.
"""
def __init__(
self,
in_channels: int,
in_shape: Tuple[int, int],
out_shape: Tuple[int, int],
grid_in: Optional[str] = "equiangular",
grid_out: Optional[str] = "equiangular",
num_heads: Optional[int] = 1,
scale: Optional[Union[torch.Tensor, float]] = None,
use_qknorm: Optional[bool] = False,
bias: Optional[bool] = True,
theta_cutoff: Optional[float] = None,
k_channels: Optional[int] = None,
out_channels: Optional[int] = None,
optimized_kernel: Optional[bool] = True,
):
if not optimized_kernels_is_available():
raise RuntimeError("Optimized kernels are required to run DistributedNeighborhoodAttentionS2.")
# initialise base class (builds global psi, creates parameters)
super().__init__(
in_channels,
in_shape,
out_shape,
grid_in=grid_in,
grid_out=grid_out,
num_heads=num_heads,
scale=scale,
use_qknorm=use_qknorm,
bias=bias,
theta_cutoff=theta_cutoff,
k_channels=k_channels,
out_channels=out_channels,
optimized_kernel=True,
)
# ---- distributed info ----
self.comm_size_polar = polar_group_size()
self.comm_rank_polar = polar_group_rank()
self.comm_size_azimuth = azimuth_group_size()
self.comm_rank_azimuth = azimuth_group_rank()
# split shapes
self.lat_in_shapes = compute_split_shapes(self.nlat_in, self.comm_size_polar)
self.lon_in_shapes = compute_split_shapes(self.nlon_in, self.comm_size_azimuth)
self.lat_out_shapes = compute_split_shapes(self.nlat_out, self.comm_size_polar)
self.lon_out_shapes = compute_split_shapes(self.nlon_out, self.comm_size_azimuth)
# local sizes for this rank
self.nlat_in_local = self.lat_in_shapes[self.comm_rank_polar]
self.nlon_in_local = self.lon_in_shapes[self.comm_rank_azimuth]
self.nlat_out_local = self.lat_out_shapes[self.comm_rank_polar]
self.nlon_out_local = self.lon_out_shapes[self.comm_rank_azimuth]
# Uniform-pscale invariant: every azimuth rank must carry the same lon pscale.
# The global divisibility check is inherited from the serial
# NeighborhoodAttentionS2.__init__, but that is not sufficient in distributed:
# if compute_split_shapes hands different ranks different local pscales
# (e.g. nlon_in=12, nlon_out=4, comm_size_azimuth=3 -> [4,4,4] vs [2,1,1]),
# the p-shift mapping in the ring exchange is ill-defined.
if self.upsample:
pscale_lon = self.nlon_out // self.nlon_in
for r, (lon_in_r, lon_out_r) in enumerate(zip(self.lon_in_shapes, self.lon_out_shapes)):
if lon_out_r != pscale_lon * lon_in_r:
raise ValueError(
f"DistributedNeighborhoodAttentionS2: inconsistent azimuth split at rank {r}: "
f"nlon_in_local={lon_in_r}, nlon_out_local={lon_out_r}. "
f"Every azimuth rank must satisfy nlon_out_local == (nlon_out // nlon_in) * nlon_in_local "
f"= {pscale_lon} * nlon_in_local. "
f"Choose (nlon_in, nlon_out, comm_size_azimuth) so that compute_split_shapes "
f"produces uniform local pscale."
)
else:
pscale_lon = self.nlon_in // self.nlon_out
for r, (lon_in_r, lon_out_r) in enumerate(zip(self.lon_in_shapes, self.lon_out_shapes)):
if lon_in_r != pscale_lon * lon_out_r:
raise ValueError(
f"DistributedNeighborhoodAttentionS2: inconsistent azimuth split at rank {r}: "
f"nlon_in_local={lon_in_r}, nlon_out_local={lon_out_r}. "
f"Every azimuth rank must satisfy nlon_in_local == (nlon_in // nlon_out) * nlon_out_local "
f"= {pscale_lon} * nlon_out_local. "
f"Choose (nlon_in, nlon_out, comm_size_azimuth) so that compute_split_shapes "
f"produces uniform local pscale."
)
# global lon/lat offsets
self.lon_in_starts = list(accumulate([0] + self.lon_in_shapes[:-1]))
self.lon_out_starts = list(accumulate([0] + self.lon_out_shapes[:-1]))
self.lat_in_starts = list(accumulate([0] + self.lat_in_shapes[:-1]))
self.lat_out_starts = list(accumulate([0] + self.lat_out_shapes[:-1]))
self.lon_lo_out = self.lon_out_starts[self.comm_rank_azimuth]
self.lat_lo_out = self.lat_out_starts[self.comm_rank_polar]
if self.upsample:
# ---- lat halo size ----
# For the scatter direction psi rows are keyed by hi, so the halo
# radius must be known BEFORE the local psi (whose rows span the
# halo-padded input range) can be built.
self.r_lat = self._compute_r_lat_upsample()
# ---- build local psi ----
# Rows are re-keyed to the halo-padded local input lat range, cols
# are filtered to the local output lat rows and the wo component is
# pre-shifted by -lon_lo_out (see _build_local_psi_upsample).
self._build_local_psi_upsample()
else:
# ---- build local psi ----
# The global psi built by the base class covers all output lat rows.
# We filter to only the rows owned by this rank and shift the wi
# component of col_idx by lon_lo_out so that the kernel can use
# local wo directly without knowing the global lon offset.
self._build_local_psi() # also precomputes self.psi_{n_long_rows,max_row_len,mid_row_len}
# ---- lat halo size ----
# Compute r_lat from the global psi: maximum |hi_global - ho_global|
# over all (ho, hi) pairs in the neighbourhood.
# Use the lat_out_lo of our polar rank to compute ho_global.
self.r_lat = self._compute_r_lat()
# -----------------------------------------------------------------------
def _build_local_psi(self):
"""Filter global psi to local output lat rows and shift col_idx wi."""
lat_lo = self.lat_lo_out
lat_hi = lat_lo + self.nlat_out_local
# global psi from the base class (built over all nlat_out rows)
col_idx_global = self.psi_col_idx # [nnz] int64
roff_global = self.psi_roff_idx # [nlat_out+1] int64
# psi_row_idx stores the sorted permutation: value is the row index.
# psi_roff_idx[ho] .. psi_roff_idx[ho+1] gives entries for row ho.
# (The row_idx buffer is the *sort order*, not the row indices directly.)
# For the distributed case we rebuild roff for the local rows only.
# Build local roff: select rows lat_lo..lat_hi-1
roff_local = roff_global[lat_lo : lat_hi + 1] - roff_global[lat_lo] # offset by first entry
# Select the corresponding col_idx entries
start = roff_global[lat_lo].item()
end = roff_global[lat_hi].item()
col_idx_local = col_idx_global[start:end].clone()
# Shift wi by pscale * lon_lo_out so the kernel can reconstruct wip from wo_local:
# col stores hi_global * nlon_in + wi_canonical. For global wo_global = lon_lo_out + wo_local,
# the target input column is (wi_canonical + pscale * wo_global) % nlon_in. The kernel evaluates
# (wi_shifted + pscale * wo_local) % nlon_in, so pre-shifting by pscale * lon_lo_out absorbs
# the rank-offset piece. pscale = 1 when nlon_in == nlon_out (same-shape case).
nlon_in = self.nlon_in
lon_lo = self.lon_lo_out
pscale = self.nlon_in // self.nlon_out
hi_global = col_idx_local // nlon_in
wi_canon = col_idx_local - hi_global * nlon_in
wi_shifted = (wi_canon + pscale * lon_lo) % nlon_in
col_idx_shifted = hi_global * nlon_in + wi_shifted
# Build sorted row_idx for local output rows (0-indexed within local range)
# Reuse the serial sort order: just re-sort by nnz per local row
nnz_per_row = (roff_local[1:] - roff_local[:-1]).cpu()
row_idx_local = torch.argsort(nnz_per_row, descending=True).to(torch.int32)
self.register_buffer("psi_col_idx_local", col_idx_shifted, persistent=False)
self.register_buffer("psi_roff_idx_local", roff_local, persistent=False)
self.register_buffer("psi_row_idx_local", row_idx_local, persistent=False)
# Precompute the CSR long/short row split once, here in the constructor,
# on the still-on-CPU local psi buffers (split_csr_rows has a CPU path).
# The split depends only on the psi sparsity geometry, which is fixed
# after init, so it is identical on every ring step / iteration. Computing
# it once keeps it off the per-step hot path (it otherwise cost a 24-byte
# D2H sync per ring step) and off any compiled forward. Stored as plain
# Python ints and threaded into the ring-step ops.
n_long_rows, max_row_len, mid_row_len = attention_kernels.split_csr_rows.default(row_idx_local, roff_local, self.nlat_out_local)
self.psi_n_long_rows = int(n_long_rows)
self.psi_max_row_len = int(max_row_len)
self.psi_mid_row_len = int(mid_row_len)
def _compute_r_lat(self) -> int:
"""Max lat halo radius needed across all polar ranks.
Computed locally from the global psi (built identically on every rank
by the base class), so no communication is required.
"""
if polar_group_size() == 1:
return 0
col_idx = self.psi_col_idx # global, all nlat_out rows
if col_idx.numel() == 0:
return 0
roff = self.psi_roff_idx
r = 0
for rank in range(self.comm_size_polar):
lat_in_lo = self.lat_in_starts[rank]
lat_in_hi = lat_in_lo + self.lat_in_shapes[rank]
lat_out_lo = self.lat_out_starts[rank]
lat_out_hi = lat_out_lo + self.lat_out_shapes[rank]
start = roff[lat_out_lo].item()
end = roff[lat_out_hi].item()
if start == end:
continue
hi = (col_idx[start:end] // self.nlon_in).long()
r_top = max(0, lat_in_lo - int(hi.min().item()))
r_bot = max(0, int(hi.max().item()) - (lat_in_hi - 1))
r = max(r, r_top, r_bot)
return r
# -----------------------------------------------------------------------
# upsample (scatter) direction helpers
# -----------------------------------------------------------------------
def _compute_r_lat_upsample(self) -> int:
"""Max lat halo radius needed across all polar ranks, upsample direction.
In the scatter psi (rows keyed by input lat hi, cols encoding output
cells), the entries relevant to a polar rank are those whose OUTPUT row
ho falls into its local output shard; the halo is then determined by how
far the corresponding INPUT rows hi reach outside its local input shard.
Computed locally from the global psi (built identically on every rank
by the base class), so no communication is required.
"""
if polar_group_size() == 1:
return 0
col_idx = self.psi_col_idx # global, rows = nlat_in, cols = ho * nlon_out + wo
if col_idx.numel() == 0:
return 0
roff = self.psi_roff_idx
# input-lat row index of every nonzero entry
nnz_per_row = roff[1:] - roff[:-1]
hi_of_nz = torch.repeat_interleave(torch.arange(self.nlat_in, dtype=torch.int64, device=col_idx.device), nnz_per_row)
ho = (col_idx // self.nlon_out).long()
r = 0
for rank in range(self.comm_size_polar):
lat_in_lo = self.lat_in_starts[rank]
lat_in_hi = lat_in_lo + self.lat_in_shapes[rank]
lat_out_lo = self.lat_out_starts[rank]
lat_out_hi = lat_out_lo + self.lat_out_shapes[rank]
mask = (ho >= lat_out_lo) & (ho < lat_out_hi)
if not bool(mask.any()):
continue
hi = hi_of_nz[mask]
r_top = max(0, lat_in_lo - int(hi.min().item()))
r_bot = max(0, int(hi.max().item()) - (lat_in_hi - 1))
r = max(r, r_top, r_bot)
return r
def _build_local_psi_upsample(self):
"""Build the local scatter psi for the upsample ring kernels.
The global psi built by the base class has rows keyed by the input lat
hi in [0, nlat_in) and cols encoding ho * nlon_out + wo_canonical on the
fine output grid (canonical at wi = 0). The local psi
* re-keys the rows to the halo-padded LOCAL input lat range
[lat_halo_start, lat_halo_start + nlat_halo); pole-padding rows
(hi outside the global grid) are empty,
* keeps only entries whose output row ho falls into the local output
shard and re-keys ho to ho_local = ho - lat_lo_out,
* pre-shifts the wo component by -lon_lo_out (mod nlon_out) so the
kernel's mapping w = (wo_shifted + pscale_out * wi_global) mod
nlon_out directly yields the LOCAL output longitude, with the
locality test w < nlon_out_local.
"""
nlon_out = self.nlon_out
lat_lo_out = self.lat_lo_out
lat_hi_out = lat_lo_out + self.nlat_out_local
lon_lo_out = self.lon_lo_out
col_idx_global = self.psi_col_idx # [nnz] int64
roff_global = self.psi_roff_idx # [nlat_in+1] int64
nlat_halo = self.nlat_in_local + 2 * self.r_lat
lat_halo_start = self.lat_in_starts[self.comm_rank_polar] - self.r_lat
# input-lat row index of every nonzero entry
nnz_per_row = roff_global[1:] - roff_global[:-1]
hi_of_nz = torch.repeat_interleave(torch.arange(self.nlat_in, dtype=torch.int64, device=col_idx_global.device), nnz_per_row)
ho = col_idx_global // nlon_out
wo = col_idx_global - ho * nlon_out
# keep entries whose input row lies in the halo-padded local range and
# whose output row is owned by this polar rank
hi_local = hi_of_nz - lat_halo_start
mask = (ho >= lat_lo_out) & (ho < lat_hi_out) & (hi_local >= 0) & (hi_local < nlat_halo)
hi_sel = hi_local[mask]
ho_sel = ho[mask] - lat_lo_out
wo_sel = (wo[mask] - lon_lo_out) % nlon_out
col_idx_local = ho_sel * nlon_out + wo_sel
# rebuild the CSR row offsets over the halo-padded local rows; masked
# selection preserves the global row-major order, so col_idx_local is
# already CSR-consistent with roff_local
counts = torch.bincount(hi_sel, minlength=nlat_halo)
roff_local = torch.zeros(nlat_halo + 1, dtype=roff_global.dtype, device=roff_global.device)
roff_local[1:] = torch.cumsum(counts, dim=0)
self.register_buffer("psi_col_idx_local", col_idx_local.contiguous(), persistent=False)
self.register_buffer("psi_roff_idx_local", roff_local.contiguous(), persistent=False)
# -----------------------------------------------------------------------
def forward(
self,
query: torch.Tensor,
key: Optional[torch.Tensor] = None,
value: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if key is None:
key = query
if value is None:
value = query
torch._check(query.dim() == 4, lambda: f"Expected 4-dimensional query tensor, got {query.dim()} dimensions")
torch._check(key.dim() == 4, lambda: f"Expected 4-dimensional key tensor, got {key.dim()} dimensions")
torch._check(value.dim() == 4, lambda: f"Expected 4-dimensional value tensor, got {value.dim()} dimensions")
torch._check(query.shape[-2] == self.nlat_out_local, lambda: f"Expected query latitudes shape[-2]=={self.nlat_out_local}, got {query.shape[-2]}")
torch._check(query.shape[-1] == self.nlon_out_local, lambda: f"Expected query longitudes shape[-1]=={self.nlon_out_local}, got {query.shape[-1]}")
torch._check(key.shape[-2] == self.nlat_in_local, lambda: f"Expected key latitudes shape[-2]=={self.nlat_in_local}, got {key.shape[-2]}")
torch._check(key.shape[-1] == self.nlon_in_local, lambda: f"Expected key longitudes shape[-1]=={self.nlon_in_local}, got {key.shape[-1]}")
torch._check(value.shape[-2] == self.nlat_in_local, lambda: f"Expected value latitudes shape[-2]=={self.nlat_in_local}, got {value.shape[-2]}")
torch._check(value.shape[-1] == self.nlon_in_local, lambda: f"Expected value longitudes shape[-1]=={self.nlon_in_local}, got {value.shape[-1]}")
# ---- 1. project to k/v/q ----
key_proj = nn.functional.conv2d(key, self.k_weights, bias=self.k_bias)
value_proj = nn.functional.conv2d(value, self.v_weights, bias=self.v_bias)
query_proj = nn.functional.conv2d(query, self.q_weights, bias=self.q_bias)
# QK normalization (must come before scale)
if self.q_norm_weights is not None:
B, C, H, W = query_proj.shape
query_proj = query_proj.reshape(B, self.num_heads, -1, H, W).permute(0, 1, 3, 4, 2)
query_proj = nn.functional.rms_norm(query_proj, normalized_shape=self.q_norm_weights.shape, weight=1 + self.q_norm_weights)
query_proj = query_proj.permute(0, 1, 4, 2, 3).reshape(B, C, H, W).contiguous()
if self.k_norm_weights is not None:
B, C, H, W = key_proj.shape
key_proj = key_proj.reshape(B, self.num_heads, -1, H, W).permute(0, 1, 3, 4, 2)
key_proj = nn.functional.rms_norm(key_proj, normalized_shape=self.k_norm_weights.shape, weight=1 + self.k_norm_weights)
key_proj = key_proj.permute(0, 1, 4, 2, 3).reshape(B, C, H, W).contiguous()
# scale after normalization
query_proj = query_proj * self.scale
# fold num_heads into batch
B, _, H, W = key_proj.shape
key_proj = key_proj.reshape(B * self.num_heads, -1, H, W)
B, _, H, W = value_proj.shape
value_proj = value_proj.reshape(B * self.num_heads, -1, H, W)
B, _, H, W = query_proj.shape
query_proj = query_proj.reshape(B * self.num_heads, -1, H, W)
# ---- 2. lat halo exchange ----
# key_proj/value_proj: [Bnh, C, H_in_local, W_in_local]
# Use differentiable halo exchange when there is an actual polar split;
# otherwise fall through to the identity (no-op).
if self.r_lat > 0 and self.comm_size_polar > 1:
key_halo = polar_halo_exchange(key_proj, self.r_lat)
value_halo = polar_halo_exchange(value_proj, self.r_lat)
else:
key_halo = key_proj
value_halo = value_proj
# global lat index of first halo row
lat_halo_start = self.lat_in_starts[self.comm_rank_polar] - self.r_lat
# ---- 3. ring attention ----
# Under autocast, cast k/v/q to the autocast dtype before .apply() —
# mirrors PyTorch's autocast-eligible-op dataflow. Upstream Linear
# projections under autocast already produce bf16, so this is usually
# a no-op; covers the case where upstream is fp32-producing.
key_halo, value_halo, query_proj = _cast_to_autocast_dtype(key_halo, value_halo, query_proj)
if self.upsample:
# Global pscale_out — the kernel must not infer this from local shapes,
# because kernel `nlon_out` is nlon_out_local which differs when az_size > 1.
pscale_out = self.nlon_out // self.nlon_in
out, _, _ = _RingNeighborhoodAttentionUpsampleFn.apply(
key_halo,
value_halo,
query_proj,
self.psi_col_idx_local,
self.psi_roff_idx_local,
self.quad_weights,
self.nlon_in,
self.nlon_out,
pscale_out,
self.lon_in_starts, # lon chunk starts for kv (same as lon_in)
self.lon_in_shapes, # lon chunk sizes for kv
lat_halo_start,
self.nlat_out_local,
self.nlon_out_local,
self.r_lat,
azimuth_group(),
self.comm_rank_azimuth,
self.comm_size_azimuth,
) # [Bnh, C_v, H_out_local, W_out_local]
else:
# Global pscale — the kernel must not infer this from local shapes,
# because kernel `nlon_out` is nlon_out_local which differs when az_size > 1.
pscale = self.nlon_in // self.nlon_out
out, _, _ = _RingNeighborhoodAttentionFn.apply(
key_halo,
value_halo,
query_proj,
self.psi_col_idx_local,
self.psi_roff_idx_local,
self.psi_row_idx_local,
self.quad_weights,
self.nlon_in,
pscale,
self.lon_in_starts, # lon chunk starts for kv (same as lon_in)
self.lon_in_shapes, # lon chunk sizes for kv
lat_halo_start,
self.nlat_out_local,
self.nlon_out_local,
self.r_lat,
azimuth_group(),
self.comm_rank_azimuth,
self.comm_size_azimuth,
self.psi_n_long_rows,
self.psi_max_row_len,
self.psi_mid_row_len,
) # [Bnh, C_v, H_out_local, W_out_local]
# unfold num_heads
B_nh, C_v, H_out, W_out = out.shape
B_orig = B_nh // self.num_heads
out = out.reshape(B_orig, self.num_heads * C_v, H_out, W_out)
# ---- 4. output projection ----
out = nn.functional.conv2d(out, self.proj_weights, bias=self.proj_bias)
return out