Spherical attention equivariance test#

Tests rotational equivariance of spherical attention modules. The core idea: applying attention then rotating should match rotating then applying attention.

Method:

  1. Rotate input signals using SO(3) transformations

  2. Apply attention module

  3. Rotate output back

  4. Compare with direct attention application

Interpolation error provides a baseline for numerical accuracy of the rotation procedure.

Setup#

import torch
import torch.nn as nn
import torch_harmonics as th
import numpy as np
from scipy.interpolate import RegularGridInterpolator

from torch_harmonics.quadrature import precompute_latitudes
from torch_harmonics.plotting import plot_sphere

Configuration#

Grid type and device selection for the tests.

grid="equiangular"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.set_grad_enabled(False)
torch.autograd.grad_mode.set_grad_enabled(mode=False)

Helper functions#

Rotation and interpolation utilities for testing equivariance on the sphere:

  • rotate_grid: applies SO(3) YZY Euler rotations to spherical coordinates

  • interpolate_to_grid: performs cubic interpolation between grids

  • equivariance_test_signal: main test function that rotates, applies attention, and computes error metrics

def cartesian_to_spherical(x, y, z):
    """Convert from cartesian to spherical coordinates (theta, phi). phi in [0, 2pi)"""
    theta = np.arccos(z)
    phi = (np.arctan2(y, x) + 2*np.pi) % (2*np.pi)
    return theta, phi

def rotate_grid(theta_grid, phi_grid, alpha, beta, gamma):
    """Apply SO(3) YZY rotation to a lat/lon grid."""
    # Convert to cartesian
    x = np.sin(theta_grid) * np.cos(phi_grid)
    y = np.sin(theta_grid) * np.sin(phi_grid)
    z = np.cos(theta_grid)
    # Stack and flatten
    coords = np.stack([x.flatten(), y.flatten(), z.flatten()])
    # Rotation matrix for YZY (passive rotation)
    # This is a placeholder, you might use scipy.spatial.transform.Rotation if needed
    from scipy.spatial.transform import Rotation as R
    rot = R.from_euler('yzy', [gamma, beta, alpha])
    rotated = rot.apply(coords.T)
    # Convert back to spherical
    x_rot, y_rot, z_rot = rotated[:,0], rotated[:,1], rotated[:,2]
    theta_rot, phi_rot = cartesian_to_spherical(x_rot, y_rot, z_rot)
    theta_rot = theta_rot.reshape(theta_grid.shape)
    phi_rot = phi_rot.reshape(phi_grid.shape)
    return theta_rot, phi_rot

def interpolate_to_grid(signal, theta_to, phi_to, method="cubic"):
    """Interpolate signal defined on (theta_from, phi_from) to (theta_to, phi_to)."""
    orig_shape = signal.shape
    # cylindrical extension to handle boundary
    signal = torch.cat([signal, signal[..., 0:1]], dim=-1)
    phi_from = np.linspace(0, 2 * np.pi, signal.shape[-1])
    theta_from = np.linspace(0, np.pi, signal.shape[-2])
    signal = signal.reshape(-1, *signal.shape[-2:]).movedim(0, -1)
    # Assume signal: [batch, nlat, nlon], theta_from, phi_from: [nlat, nlon]
    interpolator = RegularGridInterpolator(
        (theta_from, phi_from),
        signal.cpu().detach().numpy(),
        method=method,
        bounds_error=True,
        fill_value=0.
    )
    interp_points = np.stack([theta_to.flatten(), phi_to.flatten()], axis=-1)
    signal_interp = interpolator(interp_points)
    return torch.tensor(signal_interp, dtype=signal.dtype, device=signal.device).movedim(-1, 0).reshape(orig_shape)
def equivariance_test_signal(q, k, v, attention_module, theta, phi, alpha, beta, gamma, quadrature_weights):
    """Test equivariance by rotating, running attention, rotating back, and comparing."""
    # 1. Rotate input grid
    theta_r, phi_r = rotate_grid(theta, phi, alpha, beta, gamma)
    # 2. Interpolate signal to rotated grid
    q_rot = interpolate_to_grid(q, theta_r, phi_r)
    k_rot = interpolate_to_grid(k, theta_r, phi_r)
    v_rot = interpolate_to_grid(v, theta_r, phi_r)
    # 3. Apply attention module
    out_rot = attention_module(q_rot, k_rot, v_rot)
    # 4. Undo rotation: rotate grid by SO(3) inverse, interpolate back
    theta_inv, phi_inv = rotate_grid(theta, phi, -gamma, -beta, -alpha)
    out_rot_back = interpolate_to_grid(out_rot, theta_inv, phi_inv)
    # 5. Reference: apply attention module to signal directly
    out_ref = attention_module(q, k, v)
    # 6. Compute equivariance error: weighted L2 norm over the sphere
    equiv_error = ((out_rot_back - out_ref) ** 2 * quadrature_weights).sum(dim=(-1, -2)).sqrt() / (out_ref ** 2 * quadrature_weights).sum(dim=(-1, -2)).sqrt()
    # 7. Compute reference by computing interpolation error in the rotation
    out_rotrot = interpolate_to_grid(interpolate_to_grid(out_ref, theta_r, phi_r), theta_inv, phi_inv)
    interp_error = ((out_rotrot - out_ref) ** 2 * quadrature_weights).sum(dim=(-1, -2)).sqrt() / (out_ref ** 2 * quadrature_weights).sum(dim=(-1, -2)).sqrt()
    return equiv_error, interp_error

Visualization test#

Quick visual check of the rotation and interpolation pipeline. Generates a spherical harmonic signal, rotates it, and rotates it back to verify the procedure works correctly.

torch.manual_seed(333)

nlat = 65
nlon = 2*(nlat-1)

# sh coefficients
l_degree = 12
m_degree = 12

# input signal is a spherical harmonic
isht = th.InverseRealSHT(nlat, nlon, lmax=l_degree+1, mmax=m_degree+1)
# signal = torch.zeros(l_degree+1, m_degree+1, dtype=torch.complex64)
# signal[l_degree, m_degree] = 1.0
# signal = isht(signal)
signal = torch.randn(l_degree+1, m_degree+1, dtype=torch.complex64)
signal = isht(signal)

# do the rotated signal

# Define grid (equiangular example)
theta = np.linspace(0, np.pi, nlat)
phi = np.linspace(0, 2 * np.pi, nlon, endpoint=False)
theta_grid, phi_grid = np.meshgrid(theta, phi, indexing='ij')

# Example: some SO(3) rotation
alpha, beta, gamma = 0.0, np.pi/5, np.pi/2

# passively rotate the grid and interpolate
theta_r, phi_r = rotate_grid(theta_grid, phi_grid, -gamma, -beta, -alpha)
signal_rot = interpolate_to_grid(signal, theta_r, phi_r)

# passively rotate the grid back and interpolate
theta_rr, phi_rr = rotate_grid(theta_grid, phi_grid, alpha, beta, gamma)
signal_rotrot = interpolate_to_grid(signal_rot, theta_rr, phi_rr)

plot_sphere(signal, projection="orthographic", central_longitude=180, colorbar=True)
plot_sphere(signal_rot, projection="orthographic", central_longitude=180, colorbar=True)
plot_sphere(signal_rotrot, projection="orthographic", central_longitude=180, colorbar=True)
plot_sphere(signal_rotrot - signal, projection="orthographic", central_longitude=180, colorbar=True)
<cartopy.mpl.geocollection.GeoQuadMesh at 0x12f86d2e0>
../_images/f14a4de18e2c2ff4852ab673c7ca90b7a1f20da9dd7838388e139ee97638620a.png ../_images/ce4dcf4c4db9eb4439168965cd996f9d2511b27eea1337fabc98be7d8d24aeb9.png ../_images/da64c583b95bb05a21077dfb749165b347fd691f9e39b139abcc799441d16550.png ../_images/1f6820fea76587e8c499939af38c05e772de0183e52d315354961096e0640586.png
channels = out_channels = 1
k_channels = 3

attention_module = th.AttentionS2(channels, 1, (nlat, nlon), (nlat, nlon), grid_in=grid, grid_out=grid, k_channels=k_channels, out_channels=out_channels)

signal = attention_module(signal.unsqueeze(0).unsqueeze(0)).squeeze()
signal_rot = attention_module(signal_rot.unsqueeze(0).unsqueeze(0)).squeeze()

plot_sphere(signal.detach(), projection="orthographic", central_longitude=180)
plot_sphere(signal_rot.detach(), projection="orthographic", central_longitude=180)
<cartopy.mpl.geocollection.GeoQuadMesh at 0x12f97ffb0>
../_images/4fde098c5049e36fced65045fbedad6af4e7efc98dc915f8745879bcb43c3b86.png ../_images/03a8d1f7e947489beffa5e9e1ad1f0e6a8224122ffc86bacea8199c3895e83bf.png

Equivariance test: identical query, key, value#

Tests AttentionS2 with the same input for query, key, and value across multiple grid resolutions. Measures how equivariance error scales with grid refinement. Lower resolutions have higher interpolation artifacts.

torch.manual_seed(333)
torch.cuda.manual_seed(333)

in_channels = out_channels = 1
k_channels = 1

# sh coefficients
l_degree = 12
m_degree = 12

q_weights = nn.Parameter(torch.randn(k_channels, in_channels, 1, 1, device=device))
k_weights = nn.Parameter(torch.randn(k_channels, in_channels, 1, 1,  device=device))
v_weights = nn.Parameter(torch.randn(out_channels, in_channels, 1, 1,  device=device))
proj_weights = nn.Parameter(torch.randn(out_channels, out_channels, 1, 1,  device=device))

nlats = 2**np.arange(3,8) + 1


coeffs = torch.randn(1, in_channels, l_degree+1, m_degree+1, dtype=torch.complex64)

print(f"nlats={nlats}")

for nlat in nlats:

    nlon = 2*(nlat-1)

    # Define grid (equiangular example)
    theta = np.linspace(0, np.pi, nlat)
    phi = np.linspace(0, 2 * np.pi, nlon, endpoint=False)
    theta_grid, phi_grid = np.meshgrid(theta, phi, indexing='ij')

    # Assume quadrature weights are given or use Gaussian quadrature as done in torch-harmonics for accuracy
    # quadrature_weights = np.ones((nlat, nlon)) * (4*np.pi) / (nlat*nlon) # Placeholder
    _, quad_weights = precompute_latitudes(nlat, grid=grid)
    quad_weights = 2 * np.pi * quad_weights.reshape(-1, 1) / nlon

    # input signal is a spherical harmonic
    isht = th.InverseRealSHT(nlat, nlon, lmax=l_degree+1, mmax=m_degree+1)
    # signal = torch.zeros(l_degree+1, m_degree+1, dtype=torch.complex64)
    # signal[l_degree, m_degree] = 1.0
    # signal = isht(signal).to(device)
    signal = isht(coeffs).to(device)

    # Example: some SO(3) rotation
    # alpha, beta, gamma = 0, np.pi/5, 0
    alpha, beta, gamma = np.pi/3, np.pi/5, np.pi/2

    # Define your torch-harmonics spherical attention module
    attention_module = th.AttentionS2(in_channels, 1, (nlat, nlon), (nlat, nlon), grid_in=grid, grid_out=grid, k_channels=k_channels, out_channels=out_channels).to(device)
    # attention_module = th.NeighborhoodAttentionS2(in_channels, (nlat, nlon), (nlat, nlon), grid_in=grid, grid_out=grid, k_channels=k_channels, out_channels=out_channels, theta_cutoff=0.25 * np.pi, optimized_kernel=True).to(device)
    attention_module.q_weights = q_weights
    attention_module.k_weights = k_weights
    attention_module.v_weights = v_weights
    attention_module.proj_weights = proj_weights

    equiv_error, interp_error = equivariance_test_signal(
        signal, signal, signal, attention_module,
        theta_grid, phi_grid,
        alpha, beta, gamma,
        torch.tensor(quad_weights, device=device)
    )
    print(f"grid: {nlat}x{nlon}")
    print(f"Interpolation error: {interp_error.item()}")
    print(f"Equivariance error: {equiv_error.item()}")
nlats=[  9  17  33  65 129]
grid: 9x16
Interpolation error: 0.547884481270984
Equivariance error: 0.7313590599527867
grid: 17x32
Interpolation error: 0.27779018421011237
Equivariance error: 0.4753986803499113
grid: 33x64
Interpolation error: 0.1639783752757204
Equivariance error: 0.2584725718742366
grid: 65x128
Interpolation error: 0.08244463632493265
Equivariance error: 0.12208064529073714
/var/folders/zb/v9dmh1fn4pl8d1n8452xmw_h0000gp/T/ipykernel_14179/3010684959.py:60: UserWarning: To copy construct from a tensor, it is recommended to use sourceTensor.detach().clone() or sourceTensor.detach().clone().requires_grad_(True), rather than torch.tensor(sourceTensor).
  torch.tensor(quad_weights, device=device)
grid: 129x256
Interpolation error: 0.03347704813575838
Equivariance error: 0.044356914035899926

Equivariance test: distinct query, key, value#

Tests NeighborhoodAttentionS2 with different inputs for query, key, and value. Runs on a batch of random spherical harmonic signals to compute statistical measures (mean and standard deviation) of equivariance error across multiple samples.

torch.manual_seed(333)
torch.cuda.manual_seed(333)

batch_size = 32

in_channels = out_channels = 1
k_channels = 1

# sh coefficients
l_degree = 12
m_degree = 12

q_weights = nn.Parameter(torch.ones(k_channels, in_channels, 1, 1, device=device))
k_weights = nn.Parameter(torch.ones(k_channels, in_channels, 1, 1,  device=device))
v_weights = nn.Parameter(torch.ones(out_channels, in_channels, 1, 1,  device=device))
proj_weights = nn.Parameter(torch.ones(out_channels, out_channels, 1, 1,  device=device))

nlats = 3 * 2**np.arange(2,7) + 1

coeffs_q = torch.randn(batch_size, in_channels, l_degree+1, m_degree+1, dtype=torch.complex64)
coeffs_k = torch.randn(batch_size, in_channels, l_degree+1, m_degree+1, dtype=torch.complex64)
coeffs_v = torch.randn(batch_size, in_channels, l_degree+1, m_degree+1, dtype=torch.complex64)

print(f"nlats={nlats}")

for nlat in nlats:

    nlat = int(nlat)
    nlon = 2*(nlat-1)

    # Define grid (equiangular example)
    theta = np.linspace(0, np.pi, nlat)
    phi = np.linspace(0, 2 * np.pi, nlon, endpoint=False)
    theta_grid, phi_grid = np.meshgrid(theta, phi, indexing='ij')

    _, quad_weights = precompute_latitudes(nlat, grid=grid)
    quad_weights = 2 * np.pi * quad_weights.reshape(-1, 1) / nlon

    # input signal is a spherical harmonic
    isht = th.InverseRealSHT(nlat, nlon, lmax=l_degree+1, mmax=m_degree+1)
    q = isht(coeffs_q).to(device)
    k = isht(coeffs_k).to(device)
    v = isht(coeffs_v).to(device)

    alpha, beta, gamma = np.pi/3, np.pi/5, np.pi/2

    attention_module = th.NeighborhoodAttentionS2(in_channels, (nlat, nlon), (nlat, nlon), grid_in=grid, grid_out=grid, k_channels=k_channels, out_channels=out_channels, theta_cutoff=0.25 * np.pi, optimized_kernel=True).to(device)
    attention_module.q_weights = q_weights
    attention_module.k_weights = k_weights
    attention_module.v_weights = v_weights
    attention_module.proj_weights = proj_weights

    interp_errors = []
    equiv_errors = []

    for i in range(batch_size):
        equiv_error, interp_error = equivariance_test_signal(
            q[i:i+1], k[i:i+1], v[i:i+1], attention_module,
            theta_grid, phi_grid,
            alpha, beta, gamma,
            torch.as_tensor(quad_weights, device=device)
        )
        equiv_errors.append(equiv_error.sum().item())
        interp_errors.append(interp_error.sum().item())

    print(f"grid: {nlat}x{nlon}")
    std, mean = torch.std_mean(torch.Tensor(interp_errors))
    print(f"Interpolation error mean: {mean.item()} std: {std.item()}")
    std, mean = torch.std_mean(torch.Tensor(equiv_errors))
    print(f"Equivariance error mean: {mean.item()} std: {std.item()}")
nlats=[ 13  25  49  97 193]
grid: 13x24
Interpolation error mean: 0.3527798056602478 std: 0.04832783713936806
Equivariance error mean: 0.8970544338226318 std: 0.14966122806072235
grid: 25x48
Interpolation error mean: 0.20696806907653809 std: 0.019902728497982025
Equivariance error mean: 0.5265681743621826 std: 0.07344599068164825
grid: 49x96
Interpolation error mean: 0.10652338713407516 std: 0.012716345489025116
Equivariance error mean: 0.21776877343654633 std: 0.02615395560860634
grid: 97x192
Interpolation error mean: 0.04518754780292511 std: 0.00758200092241168
Equivariance error mean: 0.08621252328157425 std: 0.014278355054557323
grid: 193x384
Interpolation error mean: 0.014401651918888092 std: 0.0031418600119650364
Equivariance error mean: 0.029198382049798965 std: 0.0055313510820269585