Proving tensor indices at compile time

TL;DR — ML code is full of index arithmetic, and a wrong index rarely crashes. It reads another token's data and the model gets slightly worse. Vx can hand that arithmetic to the z3 theorem prover while it compiles. This post shows two examples: the offset of one entry in a KV cache, and a copy of attention heads into shared memory. In both, the compiler proves every index is inside the tensor, for every input, and refuses common mistakes. It also shows a case where the proof held and the program printed a negative offset, and why.

A test checks the inputs you write down. A proof checks all of them. For index arithmetic that difference matters, because the bad input is usually the last row, the last layer or the last head, and that is the one a test forgets.

Vx gives you two ways to ask for a proof. A function can state what it needs (requires) and what it promises (ensures). And the low-level raw:: loads and stores that move data between memories must be proved to stay inside their tensors before the program compiles. In both cases the compiler turns the code into equations and asks z3 whether any input breaks the promise. If z3 finds one, the build fails. If z3 shows that no such input exists, the promise holds for every allowed input, and z3 did not have to try them one by one. (Later we will see the catch: z3's numbers are not quite the program's.)

An earlier post, What comptime can do, showed a contract on a small tiled matrix. This post uses two pieces of code from a transformer, and spends more time on what the proof does not cover. Every example was compiled and run with vxc built from commit fd44f29f on main, with z3 installed.

Example 1: an offset into a KV cache

During generation a transformer keeps the keys and values of every token it has seen, so it does not recompute them. This is the KV cache. Take Llama 2 7B: 32 layers, a 4,096-token context, and 32 heads of 128 numbers each. A common layout stores its keys as one flat array, ordered by layer, then token position, then head, then the numbers inside one head:

All keys: 32 layers, 536,870,912 numbers layer 0 layer 1 layer 2 … layer 31 One layer: 4,096 positions of 4,096 numbers pos 0 pos 1 pos 2 … pos 4095 One position: 32 heads of 128 numbers head 0 head 1 head 2 … head 31 One head: 128 numbers d 0 d 1 d 2 … d 127 offset = layer×16,777,216 + pos×4,096 + head×128 + d = 16,777,216 + 4,096 + 128 + 1 = 16,781,441
Each level is a run of blocks of the level below. To reach one number, step over whole layers, then whole positions, then whole heads. The highlighted number (layer 1, position 1, head 1, number 1) is at offset 16,781,441.

The offset of one number is the sum of those steps. In Vx, with the multiplications grouped the way a program usually writes them:

// Llama 2 7B KV cache, laid out as [layer][position][head][dim].
fn kv_offset(layer: i32, pos: i32, head: i32, d: i32) -> i32
requires 0 <= layer && layer < 32 && 0 <= pos && pos < 4096
      && 0 <= head && head < 32 && 0 <= d && d < 128
ensures 0 <= return && return < 536870912 {
  let row = (layer * 4096 + pos) * 32 + head;
  return row * 128 + d;
}

fn main() -> i32 {
  print(kv_offset(31, 4095, 31, 127));
  return 0;
}

The requires line says which inputs are allowed. The ensures line promises that the answer lands inside the cache, which holds 32 × 4096 × 32 × 128 = 536,870,912 numbers. The compiler proves the promise, and the program prints the last valid offset:

536870911

That is one proof covering all 536,870,912 inputs. Now make one mistake: multiply by 33 heads instead of 32, the kind of slip you make when a model config changes under you.

  let row = (layer * 4096 + pos) * 33 + head;
Error[E8001]: Function 'kv_offset' cannot prove postcondition (ensures) at compile time

The build stops. With a stride of 33 the last entry lands at offset 553,647,999, past the end of the cache, in memory that belongs to something else. In Python or C++ this compiles, runs, and reads whatever happens to live there.

Where the proof held and the program was wrong

Now a larger model: Llama 3.1 70B, with the context capped at 32,768 tokens. It has 80 layers and 8 key-value heads of 128 numbers each. Same layout, same contract, new numbers:

fn kv_offset(layer: i32, pos: i32, head: i32, d: i32) -> i32
requires 0 <= layer && layer < 80 && 0 <= pos && pos < 32768
      && 0 <= head && head < 8 && 0 <= d && d < 128
ensures 0 <= return && return < 2684354560 {
  let row = (layer * 32768 + pos) * 8 + head;
  return row * 128 + d;
}

The compiler proves this contract too. Then kv_offset(79, 32767, 7, 127) prints:

-1610612737

The promise said the offset is never negative. The compiler proved it, and the program printed a negative offset. z3 reasons about mathematical integers, which have no largest value. An i32 holds numbers only up to 2,147,483,647, and the last offset here is 2,684,354,559. The last multiplication gives a number larger than an i32 can hold, so it overflows and wraps around to a negative number. So the proof is correct for z3's integers and wrong for the program's.

The earlier post advised bounding every input in requires so that no step can overflow. This example bounds every input, and a step still overflows, because the result is too big. The advice was incomplete: every intermediate value and the result must fit in the type as well, and today the prover does not check that for you.

The fix for this program is to use 64-bit integers. With i64 everywhere the same contract is proved, and the program prints 2684354559, the correct last offset.

The fix for the language is underway. Vx is adding a usize type for sizes and offsets (#1139), and its arithmetic already stops the program with an error when it overflows (#1149). An offset like this one would stop the program at the multiplication that overflows. The prover itself still treats numbers as unbounded, so it cannot yet warn you, while it compiles, that a 32-bit or 64-bit value will overflow.

Example 2: copying attention heads into shared memory

An attention kernel on a GPU first copies a tile of keys from the GPU's main memory into the small, fast shared memory next to its cores. In Vx you can write the code that moves a tile between two memories yourself. You implement the Transfer trait (an interface, in other languages) for that pair of memories, using raw::load and raw::store. These are low-level operations, and they skip the bounds check while the program runs. So the compiler requires a proof, before it compiles the program, that every index they use is inside its tensor.

Here is a copy of four heads of 64 numbers each, stored one head after another. The Topology block describes the GPU's memories (its L2 cache and its shared memory, SMEM), so the example compiles on its own:

Topology Dev {
  memory: Memory::L2,
  visible: [Memory::L2, Memory::SMEM],
  transfer Memory::CPU_DRAM -> Memory::L2 : 10 GB/s,
  transfer Memory::L2 -> Memory::SMEM
}
impl Transfer<Memory::L2, Memory::SMEM> for Topology::Dev {
  fn move_tile(src: &Tensor<f32, [256]>, dst: &mut Tensor<f32, [256]>) -> i32 {
    for h in 0..4 {
      for d in 0..64 {
        raw::store(dst, h * 64 + d, raw::load(src, h * 64 + d));
      }
    }
    raw::barrier();
    return 0;
  }
}

There is no contract to write. Each for loop tells the prover the range of its counter, 0 <= h < 4 and 0 <= d < 64, so it can show that h * 64 + d runs from 0 to 255 and the tensor has 256 numbers. This compiles.

Now two small mistakes. A stride of 65 instead of 64, as if each head had one extra number:

        raw::store(dst, h * 65 + d, raw::load(src, h * 64 + d));

And an index shifted by one:

        raw::store(dst, h * 64 + d + 1, raw::load(src, h * 64 + d));

Both stop the build with the same error, which names the operation, the tensor and the range it needed:

Error[E6018] at 11:9: cannot prove the index of `raw::store` stays inside `dst` (needs 0 <= index < 256).
Prove it with a loop bound or invariant, or assert it in an `unsafe` block

On a GPU, neither mistake would crash. The extra writes would land in shared memory that another part of the kernel is using, and the attention scores would be slightly wrong.

If you know something the prover cannot see, you can wrap the access in unsafe. The error then becomes a warning that says what you did:

Warning: cannot prove the index of `raw::store` stays inside `dst` (needs 0 <= index < 256);
asserted, not proven (unsafe)

What a bounds proof does not tell you

Swap the layout by mistake, so the copy writes d * 4 + h instead of h * 64 + d. That index also runs from 0 to 255. The program compiles, because every access is in bounds, and the heads come out interleaved. The prover answers the question it was asked, “is this index inside the tensor?”, and nothing more. Whether the data ends up in the right order is still a job for tests.

Other limits today

The contract syntax is described in the book's contracts chapter, and the raw:: bounds rule in docs/custom_transfer_contract.md. Vx is Apache 2.0 with the LLVM exception — install it and try changing a 64 to a 65.