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