xla 0.3.1

Bindings for the XLA C++ library.
Documentation

xla-rs

Experimentation using the xla compiler from rust

Pre-compiled binaries for the xla library can be downloaded from the elixir-nx/xla repo. These should be extracted at the root of this repository, resulting in a xla_extension subdirectory being created, the currently supported version is 0.10.0.

For a linux platform, this can be done via:

wget https://github.com/elixir-nx/xla/releases/download/v0.10.0/xla_extension-0.10.0-x86_64-linux-gnu-cpu.tar.gz
tar -xzvf xla_extension-0.10.0-x86_64-linux-gnu-cpu.tar.gz

If the xla_extension directory is not in the main project directory, the path can be specified via the XLA_EXTENSION_DIR environment variable.

Generating some Text Samples with Qwen3.5

The Qwen3.5 models can be used to generate text. They use a hybrid architecture mixing gated DeltaNet linear attention layers with full attention layers. The example loads the safetensors weights directly and runs greedy generation with a kv-cache and DeltaNet state carry-over.

The weights and tokenizer config are downloaded automatically from the hub (and cached locally). Generation runs in bf16 by default, --dtype f32 and --dtype f16 are also supported. The model size is selected with --which, one of 0.8b, 2b (the default), 4b, or 9b:

# Run the example, use --cpu to run on cpu rather than gpu.
cargo run --example qwen35 --release --features hf-hub -- \
  --which 2b \
  --prompt "What is the capital of France? Answer in one word." \
  --sample-len 30

Comparison with transformers

Greedy generation of 60 tokens from a 16 token prompt, measured after compilation and weight loading, on a RTX 4080 SUPER (16GB) for the GPU rows and a Ryzen 9 7950X (16 cores) for the CPU rows. Prefill is the time to the first generated token, the decode rate covers the remaining 59 tokens. GPU numbers are the median of 3 runs, the spread mostly comes from the gemm autotuner (see the section below). The generated tokens are identical between the two implementations for all the configurations below, except 0.8B on GPU bf16 where transformers diverges after 9 tokens because of bf16 rounding (the f32 outputs of both implementations agree with the xla-rs bf16 tokens).

Decode rate:

0.8B 2B 4B
GPU bf16, xla-rs 341.8 tok/s 153.7 tok/s 73.3 tok/s
GPU bf16, transformers 73.1 tok/s 68.8 tok/s 51.1 tok/s
CPU f32, xla-rs 15.7 tok/s 6.7 tok/s -
CPU f32, transformers 7.9 tok/s 3.9 tok/s -

Prefill (time to first token):

0.8B 2B 4B
GPU bf16, xla-rs 56 ms 61 ms 77 ms
GPU bf16, transformers 57 ms 58 ms 84 ms
CPU f32, xla-rs 348 ms 557 ms -
CPU f32, transformers 206 ms 344 ms -

Note that the xla-rs prefill always processes the full padded 128 token context (with a sequential scan for the DeltaNet layers), while transformers only processes the 16 prompt tokens, which explains the slower xla-rs prefill on cpu. The 4B model in f32 does not fit in the 32GB of RAM of the benchmark machine. Versions: xla-rs with xla_extension 0.10.0 (CUDA 13.0 build), transformers 5.14.1 with torch 2.13.0 (cu130 on gpu, cpu wheel on cpu). The transformers numbers use the plain torch DeltaNet path, the fused flash-linear-attention kernels are not installed.

Pinning the gpu autotune configuration

The XLA gemm autotuner benchmarks candidate kernel configurations during compilation and its choices vary from compile to compile, which moves the gpu decode rate by up to ~10% on the larger models (63 to 73 tok/s across compiles of the 4B model). This is the main source of run-to-run variance in the benchmark above. The autotuner choices can be pinned to a file with the --autotune-cache flag: when the file does not exist the results of the compilation are written to it, when it exists they are reused, making the kernel selection deterministic and speeding up the compilation (~14s down to ~4.5s for the 4B model) as all the gemm tuning is skipped:

cargo run --example qwen35 --release --features hf-hub -- \
  --which 4b --autotune-cache qwen35-4b.pbtxt \
  --prompt "What is the capital of France? Answer in one word."

Pinning reproduces whatever run produced the file, good or bad: if the run that wrote the cache showed a low decode rate, delete the file and rerun until the tuning run is fast, then keep that file. The cache is keyed by fusion, gpu, and xla version, so use one file per model size and regenerate it after changing the model code or the xla_extension build (stale entries are silently re-tuned, which brings the nondeterminism back). Two smaller things to be aware of: when loading, the first prefill pays ~130ms of one-time kernel loading that is otherwise hidden inside the autotuning phase, and --xla_gpu_exhaustive_tiling_search is not a good alternative as it selects kernels on isolated micro-benchmarks and ends up slower end-to-end.

In the library this is exposed as PjRtClient::compile_with_autotune_cache(computation, load_from, dump_to) which scopes the xla_gpu_load_autotune_results_from and xla_gpu_dump_autotune_results_to debug options to a single compilation (the same options can also be set process-wide through the XLA_FLAGS environment variable).

Gemma 4 E2B and E4B

The gemma4 example runs the text model of Gemma 4 E2B or E4B (--which e2b|e4b), MatFormer-style hybrids where most decoder layers use sliding-window attention, the last layers reuse the k/v states of earlier layers, and a per-layer embedding (PLE) is mixed into the residual stream after each mlp. Generation runs in bf16 with the same prefill/decode split and on-device kv-cache as the qwen35 example, and the --autotune-cache flag is supported too.

The PLE table is the "effective" in the E model names: it holds roughly half of the parameters (2.4B of E2B's 4.6B, 2.9B of E4B's 7.5B) but is only ever read one row per token, so the example keeps it in host memory and gathers the rows for the current tokens on the cpu, passing them to the computations as an input. This halves the device memory (4.5GB of weights for E2B, 9.2GB for E4B) and is what lets E4B run on a 16GB gpu in bf16, at the cost of a host round-trip per decode step (a few percent decode rate on E2B).

The repositories are gated on the hub: accept the license on the model page and make a token available via HF_TOKEN or ~/.cache/huggingface/token for the initial download.

cargo run --example gemma4 --release --features hf-hub -- \
  --which e4b \
  --prompt "What is the capital of France? Answer in one word."

Comparison with transformers

Greedy generation of 106 tokens from a 19 token prompt, measured after compilation, weight loading, and the one-time kernel loading of the first execution, on the same RTX 4080 SUPER (16GB) and with the same versions as the qwen35 benchmark above. Prefill is the time to the first generated token (median of 5), the decode rate covers the remaining 105 tokens (median of 3 runs), and the memory column is the peak gpu usage of the xla-rs run:

E2B E4B
prefill, xla-rs 13 ms 24 ms
prefill, transformers 25 ms -
decode, xla-rs 110.1 tok/s 48.0 tok/s
decode, transformers 46.1 tok/s -
gpu memory, xla-rs 8.5 GB 14.7 GB

The transformers E4B column is empty as the full bf16 text model (15GB of weights with the PLE table on the device) does not leave enough room on a 16GB gpu; splitting it across gpu and cpu with accelerate works for checking outputs but is not a meaningful speed comparison. The generated tokens are identical between the two implementations on the prompts tested (E4B was checked against a gpu+cpu split run). One caveat: on E2B the very first token of the benchmark prompt is a near-tie (the top-2 logits are 0.11 apart in a f32 reference run, about one bf16 ulp at that magnitude), so bf16-level numeric changes such as a different autotuned gemm kernel or the PLE offload's slightly different graph can flip it and the continuations then diverge while staying equally valid.