ZML

Last month, I wrapped up my internship at ZML.

ZML is building a framework and an inference server, with the bet that pairing the compiler and hardware lets you specialize inference end-to-end, allowing better per-vendor performance without compromising with model-specific conditionals.

I worked on adding a new compiler target: the Vulkan Graphics API, to our heterogeneous inference stack.

My work on generating compute shaders for LLM inference was a long and arduous process, and I’d like to share it here.

Graph Compilers

This is ZML’s execution model:

  1. Given a model runtime in Zig, ZML represents the computation graph as symbolic tensors
  2. ZML uses the XLA graph compiler to optimize the computation graph.
  3. Execute that optimized program on the actual hardware, in the respective hardware’s assembly language. In OpenXLA, the boundary between 2 and 3 is PJRT: a contract specifying everything our framework needs to dispatch work to a backend, abstracting away accelerator specifics. PJRT plugins are concrete implementations of this contract.

Each plugin consists of a compiler + runtime. For example, given a NVIDIA GPU, a client can select the PJRT XLA CUDA plugin, with XLA GPU + CUDA runtime. Alternatively, they may choose to use the IREE CUDA plugin. Or they may choose the IREE Vulkan plugin.

ZML uses PJRT to add plugins you wish to install, and to add Vulkan, I had to write one of them.

A bit of context about Vulkan if you haven’t heard of it: Vulkan is a cross-platform graphics and compute API designed for high-performance graphics. It runs on nearly all modern GPUs from vendors including NVIDIA, AMD, Intel, ARM, and Qualcomm.

Naturally, XLA already has an emission pipeline for GPUs. So, all I had to do was extend the existing XLA:GPU pipeline to output SPIR-V instructions.

If only it were that easy…

What Vulkan expects from a compute shader

Vulkan is quite an explicit low-level API where every compute dispatch, memory barrier, and queue submission requires direct host-side driver interaction.

Unlike higher-level or more mature backends that might handle scheduling or batching under the hood, XLA's default compilation pipeline initially lowers graph operations into individual dispatches and explicit synchronization points.

Now, here’s a sample of what Vulkan expects. Instead of a functional approach like:

output = matmul(a, b);

We instead have something closer to:

// oversimplified Vulkan syntax 

compute_shader(
    binding_0 = a_tensor,
    binding_1 = b_tensor,
    binding_2 = output_tensor
);

For each variable, Vulkan defines a resource slot inside a descriptor set layout. Slots are populated with actual GPU buffers later. This feels somewhat analogous to ZML tensors being symbolic first and then populated with concrete buffers later.

Shader code        = function body
Descriptor layout  = function signature
Descriptor set     = actual arguments
vkCmdDispatch      = function call over many parallel threads

A typical ML operation reads input tensors, reads weights, and writes output tensors. The descriptor set layout could then look like:

vk::DescriptorSetLayoutBinding bindings[] = {
    // Binding 0: Input tensor
    {
        .binding = 0,
        .descriptorType = vk::DescriptorType::eStorageBuffer,
        .descriptorCount = 1,
        .stageFlags = vk::ShaderStageFlagBits::eCompute
    },

    // Binding 1: Weights
    {
        .binding = 1,
        .descriptorType = vk::DescriptorType::eStorageBuffer,
        .descriptorCount = 1,
        .stageFlags = vk::ShaderStageFlagBits::eCompute
    },

    // Binding 2: Output tensor
    {
        .binding = 2,
        .descriptorType = vk::DescriptorType::eStorageBuffer,
        .descriptorCount = 1,
        .stageFlags = vk::ShaderStageFlagBits::eCompute
    }
};

Then, on the shader side, resources are declared using layout() with {set, binding} coordinates.

For example:

layout(set = 0, binding = 0) buffer Input {
    float input[];
};

layout(set = 0, binding = 1) buffer Weights {
    float weights[];
};

layout(set = 0, binding = 2) buffer Output {
    float output[];
};

void main() {
    // read input[]
    // read weights[]
    // write output[]
}

So the rough mental model is that we separate the parameters from the kernels, they are stored collection/group of resources, with the logical mappings “binding” the resources to the groups at the top of a compute shader.

My Vulkan implementation underneath PJRT did exactly that: binding descriptors, creating buffers and pipelines, recording command buffers, inserting synchronization barriers, dispatching compute work, and submitting command buffers to a Vulkan queue.

The hardest bug

The hardest technical challenge I solved was a silent predicate-layout mismatch in XLA’s Vulkan backend.

User: Write a story

Assistant: I can’t assist with writing, but I can’t fulfill your request to write a story

It was especially difficult because nothing was actually crashing.

First of all, according to Vulkan validation tools, the IR was fully legal. Second of all, buffers were written with mostly intelligible text. Last, only certain predicate patterns exposed the bad stride clearly.

I bisected to the first incorrect activation. It was a causal mask in the attention layer that should have been lower-triangular, but had some wrong predicates.

However, even this was downstream of the real issue! The actual bug occurred much earlier, where LLVM memory operations were lowered to Vulkan descriptor accesses.

XLA made a reasonable but incorrect assertion: since SPIR-V storage buffers cannot directly represent LLVM i1 values as addressable elements, XLA predicates were exposed to Vulkan as a single byte per boolean. So it widened predicate elements from i1 to i8.

This worked for all external predicate-buffer ABIs, but it was also applied to an LLVM <N x i1> memory operation (the aforementioned causal select).

So LLVM was acting under the assumption that we had vector-packed bits (one bit = 1 lane), while the normalized Vulkan access treated it as <N x i8> (one byte = 1 lane).

As a result, the shader read predicate lane n from byte n, even though LLVM had stored that lane as bit n and thus masks and selections consumed truth values from the wrong locations.

This was a case where the IR was semantically reasonable at the value level, but incorrect at the memory/ABI boundary.

I fixed it in our Vulkan buffer-access normalization pass:

  1. Replace vector i1 loads with individual i8 loads,
  2. Truncate each byte to i1 and reconstruct the vector,
  3. Lower vector stores into individual i8 stores,
  4. Zero-extend each i1 to canonical 0/1, After that, the transformer prefill layer compiled successfully.

What’s interesting is that <N x i1> was fine as an SSA value, but not as a direct memory access against an ABI that promised N bytes.

Bugs like this one, which don’t cause crashes or panics, are very elusive.

A brief look at debugging execution.

One last story: once the compiler produced valid Vulkan SPIR-V codegen, I spent my time on performance.

Running Llama3 end-to-end for the first time took over two minutes to return a single token (prefill).

When I profiled the kernels, the first question was what kind of bottleneck this is.

My process is to search for the dominant bottleneck and then pick hypotheses from the matching bucket. Test, repeat.

  • If it's memory, reduce traffic or improve locality.
  • If it's compute, inspect the generated code or kernel efficiency.
  • If it's sync-bound, look for unnecessary barriers, overlaps or scheduling issues.
  • If it's transfers, reduce needless host device copies. In this case, memory traffic looked reasonable; the kernels themselves were small, and there were no obvious synchronization hotspots.

Looking closer with Nsight, the timeline showed thousands of tiny GPU dispatches and constant queue submissions, creating massive overhead and keeping the GPU starved for work. You can see in the image below:

This pointed me back to the host.

Since I hadn’t implemented command-buffer batching or more aggressive kernel fusion yet, every small operation had become its own submission.

The CPU was effectively drip-feeding the GPU thousands of tiny dispatches, so the actual bottleneck was launch and submission overhead.

Tuning performance in Vulkan required careful management of host-side dispatching, synchronization, and buffer transfers.

Among many other fixes, addressing these low-level API overheads enabled full utilization of Vulkan compute shaders for ZML’s workloads!

See below: Llama3 1B and MNIST running on Vulkan + NVIDIA 4090