How to Switch Between TPU and GPU Attention in WeatherNext 2
Set the attention_type field to 'splash_mha' for TPU splash attention, 'mha' for generic GPU-compatible dense attention, or 'triblockdiag_mha' for block-diagonal TPU-optimized attention.
WeatherNext 2 implements a configurable sparse transformer that abstracts hardware-specific optimizations behind a single configuration flag. The attention_type parameter in weathernext.utils.sparse_transformer.Transformer controls which kernel executes at runtime, allowing identical model code to run on TPU, GPU, or CPU without modification.
Attention Implementations in WeatherNext 2
The transformer supports three attention kernels with distinct hardware targets and performance characteristics:
| Attention Type | Implementation | Hardware Target | Configuration Value |
|---|---|---|---|
| Splash attention | jax.experimental.pallas.ops.tpu.splash_attention |
TPU only | 'splash_mha' |
| Dense MHA | Standard JAX multi_head_attention |
GPU, CPU, TPU (generic) | 'mha' |
| Block-diagonal MHA | Custom block-diagonal dense kernel | TPU-optimized, GPU compatible | 'triblockdiag_mha' |
Where the Switch Happens in the Code
The dispatch logic resides in Block.__call__ at lines 456–480 of utils/sparse_transformer.py. Here's the branching structure:
# From weathernext/utils/sparse_transformer.py, lines 456-480
if self.attention_type == 'triblockdiag_mha':
# Pads input into blocks, builds block-diagonal mask
# Lines 456-462: triblockdiag branch
...
elif self.attention_type == 'splash_mha':
# TPU-specific splash kernel from jax.experimental.pallas.ops.tpu
# Lines 477-480: splash branch
...
elif self.attention_type == 'mha':
# Standard dense attention, CSR mask converted to dense jnp.array
# Lines 474-476: dense MHA branch
...
The Transformer.__init__ constructor (lines 534–566) prepares hardware-specific masks and padding strategies based on this selection.
TPU Splash Attention Configuration
Splash attention delivers maximum throughput on TPU but imposes strict requirements:
from weathernext.utils.sparse_transformer import Transformer
import jax.numpy as jnp
import scipy.sparse as sp
adj_mat = sp.csr_matrix(...) # your mesh adjacency (N, N)
tpu_transformer = Transformer(
adj_mat=adj_mat,
attention_k_hop=2,
attention_type='splash_mha', # ← TPU-only splash kernel
mask_type='lazy', # memory-efficient lazy masking
num_heads=8,
# All block sizes must divide head_dim evenly
# Head dim assertion at line 307 enforces multiple of 128
block_q=128,
block_kv=128,
block_q_dkv=128,
block_kv_dkv=128,
block_q_dkv_compute=128,
block_kv_dkv_compute=128,
)
out = tpu_transformer(
node_features=jnp.ones((4, adj_mat.shape[0], 64)),
global_norm_conditioning=jnp.zeros((4, 10))
)
Critical constraint: The head dimension must be a multiple of 128. Line 307 of sparse_transformer.py enforces this with an explicit assertion. Attempting to run splash_mha on GPU raises an import error—the splash ops live exclusively under jax.experimental.pallas.ops.tpu.
GPU-Compatible Dense Attention
For GPU deployment, use standard dense multi-head attention:
from weathernext.utils.sparse_transformer import Transformer
import jax.numpy as jnp
import scipy.sparse as sp
gpu_transformer = Transformer(
adj_mat=adj_mat,
attention_k_hop=2,
attention_type='mha', # ← GPU-compatible dense attention
mask_type='full', # full dense mask
num_heads=8,
# No block size constraints required
)
out = gpu_transformer(
node_features=jnp.ones((4, adj_mat.shape[0], 64)),
global_norm_conditioning=jnp.zeros((4, 10))
)
This branch converts the sparse CSR adjacency matrix to a dense jnp.array with no special padding (lines 560–566).
Block-Diagonal Alternative for TPU
The triblockdiag MHA provides a middle ground—TPU-optimized without splash kernel dependencies:
block_transformer = Transformer(
adj_mat=adj_mat,
attention_k_hop=2,
attention_type='triblockdiag_mha', # block-diagonal padding
mask_type='full',
num_heads=8,
)
This implementation pads the node set to a multiple of mask_block_size and constructs a block-diagonal attention mask (lines 534–543). It runs on both TPU and GPU but is tuned for TPU's high-throughput block operations.
Key Configuration Files
utils/sparse_transformer.py— Core implementation with dispatch logic atBlock.__call__(lines 456–480) and mask construction inTransformer.__init__(lines 534–566)weathernext2/architecture.py— High-level model wiring showing howTransformerreceives its configurationweathernext1_gen/denoiser.py— Reference implementation defaulting toattention_type='splash_mha'(line 137)utils/mesh_transformer.py—MeshSparseTransformerwrapper that also respectsattention_type(line 56)
Summary
- Three attention types control hardware targeting:
'splash_mha'(TPU only),'mha'(GPU/CPU),'triblockdiag_mha'(TPU-optimized, GPU-compatible) - Single configuration change—set
attention_typeinTransformer—switches implementations without modifying downstream code - TPU splash attention requires head dimension multiple of 128 and TPU hardware; fails on GPU with import error
- GPU deployment uses
'mha'for broad compatibility or'triblockdiag_mha'for potential TPU migration
Frequently Asked Questions
What happens if I try to run splash attention on a GPU?
You'll encounter an import error. The splash kernel imports from jax.experimental.pallas.ops.tpu, which only exists in TPU-enabled JAX builds. Switch to attention_type='mha' or attention_type='triblockdiag_mha' for GPU execution.
Why does splash attention require block sizes that divide 128?
The TPU splash kernel uses 128-element vectorized operations internally. The assertion at line 307 of sparse_transformer.py validates that head_dim % (block_q or 128) == 0 to ensure aligned memory access patterns for optimal XLA compilation.
Can I use the same checkpoint with different attention types?
Generally no. The three implementations apply different padding strategies and mask shapes (triblockdiag pads nodes to block boundaries, splash computes separate Q/KV paddings, mha uses no padding). These change the effective tensor shapes, making checkpoints incompatible across attention types.
Does block-diagonal attention work well on GPU?
It runs correctly but without performance benefits. The triblockdiag implementation is optimized for TPU's high-throughput matrix block operations. On GPU, standard dense MHA (attention_type='mha') typically performs better due to mature cuDNN kernels.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →