# What Are the Custom CUDA/TPU Kernels in Levanter? A Deep Dive into the Kernel Library

> Explore Levanter's custom CUDA and TPU kernels for GPUs and TPUs. Discover high-performance primitives like DeepPE transport, fused cross-entropy, and SSD attention for accelerated deep learning.

- Repository: [The Marin Project/marin](https://github.com/marin-community/marin)
- Tags: deep-dive
- Published: 2026-09-10

---

**Levanter provides hand-written CUDA kernels for NVIDIA GPUs and Pallas XLA kernels for Google TPUs, located in `lib/levanter/src/levanter/kernels/`, implementing high-performance primitives like DeepPE transport, fused cross-entropy, and SSD attention.**

Levanter is a JAX-based framework for training large language models within the `marin-community/marin` repository. To maximize throughput on modern accelerators, Levantor bypasses standard JAX primitives in performance-critical paths by shipping **custom CUDA/TPU kernels** written in NVidia CUDA C++ and the Pallas XLA DSL.

## CUDA Kernels for GPU Acceleration

Levanter’s CUDA implementation focuses on the DeepPE (parallel-execution-pipeline) communication primitives and layout optimization routines, all residing under `lib/levanter/src/levanter/kernels/deepep/csrc/`.

### DeepPE Transport and Layout Kernels

The core GPU kernels implement tensor transport and memory movement for distributed training:

- **`deepep_transport_ffi.cu`** – Implements the transport layer for the DeepPE "PEP" primitive, handling inter-device communication patterns.
- **`deepep_layout_ffi.cu`** – Provides layout-aware copy kernels used when moving tensors between devices with different memory layouts.
- **`deepep_launch_compat.cuh`** – A header-only utility containing launch-configuration helpers for grid-size, block-size, and SM-architecture compatibility checks.

### Compilation and FFI Wrapper

Unlike pre-compiled binaries, Levanter compiles these kernels dynamically at runtime to match your specific GPU architecture. The [`transport_ffi.py`](https://github.com/marin-community/marin/blob/main/transport_ffi.py) wrapper handles this process:

1. **`_prepared_cuda_sources()`** builds a temporary build directory and copies the `.cu` sources.
2. The system `nvcc` compiler is invoked with architecture-specific flags (selected via `deepep_cuda_arch_flag()` such as `-arch=sm_70` or `-arch=sm_80`).
3. **`cffi.FFI().dlopen()`** loads the resulting shared object, exposing symbols like `launch_transport` and `layout_copy` to Python.

This JIT compilation ensures optimal instruction selection for your specific SM version while maintaining compatibility across different GPU generations.

## TPU Kernels via Pallas

For Google Cloud TPU execution, Levanter uses **Pallas**, a Python DSL for writing custom XLA programs that compile to TPU-specific instructions. These kernels live in `lib/levanter/src/levanter/kernels/pallas/` and are pure Python implementations that XLA compiles at runtime.

### Pallas Kernel Categories

The TPU kernel library contains several specialized sub-packages:

- **`ssd/`** – Implements "sparse-sharding-dense" operations for the SSD-based attention mechanism, optimizing memory access patterns for long sequences.
- **`short_conv/`** – Fuses multiple small matrix-multiplications into a single XLA program, reducing dispatch overhead and memory bandwidth.
- **`mamba3/`** – Contains the Mamba-3 sequence model implementation optimized for TPU through Pallas primitives.
- **`fused_cross_entropy_loss/`** – A memory-efficient cross-entropy loss that avoids intermediate materialization of large logit tensors.
- **[`splash_attention.py`](https://github.com/marin-community/marin/blob/main/splash_attention.py)** – A reference implementation of the Splash attention algorithm written entirely in Pallas.

### Kernel API Structure

Each kernel sub-package exposes a high-level Python API through [`api.py`](https://github.com/marin-community/marin/blob/main/api.py) files. These modules use **`pallas_call`** and **`pallas_program`** to construct XLA computation graphs dynamically. When you invoke a function like `ssd_attention` or `short_conv_forward`, the Pallas compiler generates a device-specific binary optimized for your TPU topology.

## Kernel Registration and Backend Dispatch

Levanter abstracts hardware differences through a unified dispatch layer that seamlessly routes operations to the appropriate kernel implementation.

### Operator Registration

In [`lib/levanter/src/levanter/kernels/__init__.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/kernels/__init__.py), the library registers primitive ops under the `levanter_kernels` namespace. These registrations point either to the compiled CUDA shared object (for GPU hosts) or to the Pallas XLA programs (for TPU hosts), presenting a single Python interface regardless of hardware.

### Runtime Dispatch Logic

When you call a kernel function like `levanter.kernels.fused_cross_entropy_loss`, the framework inspects **`jax.default_backend()`**:

- If the backend returns `'gpu'`, the call routes to the CUDA FFI functions exposed by `transport_ffi`.
- If the backend returns `'tpu'`, the call invokes the Pallas XLA program.

**Fallback behavior**: If `nvcc` is missing, the SM architecture is unsupported, or the Pallas compiler encounters an error, Levanter automatically falls back to a pure-JAX implementation. This guarantees numerical correctness at the cost of reduced performance.

## Usage Examples

The following examples demonstrate how to invoke these kernels in practice:

```python

# Using the fused cross-entropy loss on GPU

import jax
import jax.numpy as jnp
from levanter.kernels.fused_cross_entropy_loss import fused_cross_entropy

logits = jnp.random.randn(8, 50257)   # [batch, vocab]

labels = jnp.arange(8)                # dummy labels

loss = fused_cross_entropy(logits, labels)   # dispatches to CUDA kernel

print(loss)

```

```python

# Running SSD-based attention on TPU

import jax
from levanter.kernels.pallas.ssd import ssd_attention

# query, key, value are sharded JAX arrays on a TPU mesh

output = ssd_attention(query, key, value)   # uses Pallas XLA program

```

## Summary

- **GPU kernels** are hand-written CUDA files in `lib/levanter/src/levanter/kernels/deepep/csrc/`, compiled on-the-fly via [`transport_ffi.py`](https://github.com/marin-community/marin/blob/main/transport_ffi.py) using `cffi` and `nvcc`.
- **TPU kernels** are Pallas DSL programs in `lib/levanter/src/levanter/kernels/pallas/`, including SSD attention, short convolution, and fused cross-entropy implementations.
- **Unified dispatch** occurs through [`lib/levanter/src/levanter/kernels/__init__.py`](https://github.com/marin-community/marin/blob/main/lib/levanter/src/levanter/kernels/__init__.py), which selects backends based on `jax.default_backend()` with automatic pure-JAX fallback.
- **Key files** include `deepep_transport_ffi.cu` for DeepPE transport, [`splash_attention.py`](https://github.com/marin-community/marin/blob/main/splash_attention.py) for TPU attention, and the respective [`api.py`](https://github.com/marin-community/marin/blob/main/api.py) modules exposing high-level interfaces.

## Frequently Asked Questions

### Where are the CUDA kernel sources located in Levanter?

The CUDA sources reside in `lib/levanter/src/levanter/kernels/deepep/csrc/`. This directory contains `deepep_transport_ffi.cu` for transport primitives, `deepep_layout_ffi.cu` for layout-aware copies, and `deepep_launch_compat.cuh` for launch configuration utilities.

### How does Levanter compile CUDA kernels at runtime?

Levanter uses the `_prepared_cuda_sources()` function in [`transport_ffi.py`](https://github.com/marin-community/marin/blob/main/transport_ffi.py) to create a temporary build directory, copy `.cu` files, and invoke `nvcc` with architecture flags from `deepep_cuda_arch_flag()`. The resulting shared object is loaded via `cffi.FFI().dlopen()` to expose kernel functions to Python.

### What happens if a custom kernel fails to compile?

If compilation fails due to missing `nvcc`, unsupported SM versions, or other errors, Levanter automatically falls back to a pure-JAX implementation of the operation. This ensures training scripts remain functional across different hardware configurations, though potentially with reduced performance.

### Are the TPU kernels compatible with standard JAX operations?

Yes. The Pallas kernels in `lib/levanter/src/levanter/kernels/pallas/` return standard JAX arrays and integrate with the JAX ecosystem. They are compiled to XLA HLO at runtime and compose naturally with `jax.jit`, `jax.grad`, and other JAX transformations.