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:
Rotate input signals using SO(3) transformations
Apply attention module
Rotate output back
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 coordinatesinterpolate_to_grid: performs cubic interpolation between gridsequivariance_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>
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>
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