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

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 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 – 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 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, 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:


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

# 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 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, 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 for TPU attention, and the respective 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →