# How to Switch Between TPU and GPU Attention in WeatherNext 2

> Learn how to switch between TPU and GPU attention in WeatherNext 2 by setting the attention_type field to splash_mha mha or triblockdiag_mha for optimized performance.

- Repository: [Google DeepMind/weathernext](https://github.com/google-deepmind/weathernext)
- Tags: how-to-guide
- Published: 2026-08-12

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/utils/sparse_transformer.py). Here's the branching structure:

```python

# 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:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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**:

```python
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:

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/utils/sparse_transformer.py)** — Core implementation with dispatch logic at `Block.__call__` (lines 456–480) and mask construction in `Transformer.__init__` (lines 534–566)
- **[`weathernext2/architecture.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext2/architecture.py)** — High-level model wiring showing how `Transformer` receives its configuration
- **[`weathernext1_gen/denoiser.py`](https://github.com/google-deepmind/weathernext/blob/main/weathernext1_gen/denoiser.py)** — Reference implementation defaulting to `attention_type='splash_mha'` (line 137)
- **[`utils/mesh_transformer.py`](https://github.com/google-deepmind/weathernext/blob/main/utils/mesh_transformer.py)** — `MeshSparseTransformer` wrapper that also respects `attention_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_type` in `Transformer`—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`](https://github.com/google-deepmind/weathernext/blob/main/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.