Source code for torch_harmonics.distributed.distributed_attention

# 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