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:
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
- Callers are not checked against
requires. The function proves its promise assuming its inputs are allowed, but a call such askv_offset(40, 0, 0, 0), for layer 40 of a 32-layer model, compiles and prints 671088640, past the end of the cache. - Only linear arithmetic. The prover handles
+,-, comparisons, and multiplying by a constant. A product of two variables, such as tile bytesbm * bk * 4 <= 32768, cannot be proved even when it is true; the build fails withE8001. So a limit of the prover shows up as an error, and the code does not compile. - Integers are unbounded, as the 70B example shows.
- z3 must be installed to compile code that needs a proof.
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.