A couple of years ago, @jxmnop talked me
into writing a collection of
tensor puzzles.
Over time, they have become popular job interview prep
for wannabe MsTS. The puzzles give the user a specification
written in easy-to-read Python and ask them to
generate a vectorized implementation.
defloop_cumsum(a):
out = jnp.zeros_like(a)
total =0.for i inrange(a.shape[0]):
total = total + a[i]
out = out.at[i].set(total)return out
@jax.jit
defpuzzle_cumsum(a):
i = arange(a.shape[0])return jnp.sum(jnp.where(i[:,None]>= i[None,:], a[None,:],0.), axis=1)
Recently, I have been thinking about formal verification and
envisioning what it might look like for everyday coders. This seemed like
a good place to start.
In this blog, we walk through how to prove puzzle solutions correct
and build a library for proving properties of arbitrary JAX code: srush/jax-lean.
It was inspired by conversations with pmocz, nielstron, and these projects:
The code and proofs are
written by AI; the writing is human. This is a follow-up to Lean-Verified Transformers.
Background: JAX IR
JAX is a numerical transformation system
that utilizes a nearly one-to-one NumPy interface. While the primary use of JAX is machine
learning, the core framework has remained impressively ML-agnostic. Instead, it focuses on
a core set of transformations that can be reduced to simple numerical operations.
JAX is also a magic trick. The user thinks they are writing Python, while actually
it is pulling out code using a minimal core of numerical operations. Proving properties
of Python is nearly impossible, but proofs using this intermediate representation
are more tractable.
If you are interested in the details of how JAX works, they have an amazing tutorial
known as Autodidax.
Here, we will just briefly outline the process.
Let's begin with a simple function that adds two vectors and multiplies the result by 2.
defadd_then_scale(a, b):
total = a + b
return total *2.0
When we run this Python code, JAX is tracing the path and producing its
own representation of what operations were called.
a = jnp.array([1.,2.,3.], dtype=jnp.float32)
b = jnp.array([4.,5.,6.], dtype=jnp.float32)print(jax.make_jaxpr(add_then_scale)(a, b))
The intermediate representation is known as Jaxpr. It lives
between your Python code and the final lower-level representations that execute
the mathematical operations on hardware.
One day, we hope to have specs all the way down to the GPU. However, we are going
to focus on this Jaxpr level and assume everything below roughly follows what it says.
This is a horrible assumption for ML code, as numerical precision is of huge importance.
Hopefully, we can cover this in future blogs.
Let's look at the Jaxpr output from our function. It should look familiar. It takes two arguments: float32 vectors of size 3.
It adds them into a temporary variable, multiplies it by 2, and returns the result.
{ lambda ; a:f32[3] b:f32[3]. let
c:f32[3] = add a b
d:f32[3] = mul c 2.0:f32[]
in (d,) }
Of course, this code will get more complex with more realistic examples.
There is a set of operations that specify the core numerical transformations
needed. The core idea, though, is that it will remain a sequence of these operations,
without loops or other control flow.
defeye(n):
i = arange(n)return jnp.where(i[:,None]== i[None,:],1.,0.)
{ lambda ; . let
a:f32[3] = iota[dimension=0 dtype=float32 shape=(3,) sharding=None]
b:f32[3,1] = broadcast_in_dim[
broadcast_dimensions=(0,)
shape=(3, 1)
sharding=None
] a
c:f32[1,3] = broadcast_in_dim[
broadcast_dimensions=(1,)
shape=(1, 3)
sharding=None
] a
d:bool[3,3] = eq b c
e:f32[3,3] = jit[
name=_where
jaxpr={ lambda ; d:bool[3,3] f:f32[] g:f32[]. let
h:f32[3,3] = broadcast_in_dim[
broadcast_dimensions=()
shape=(3, 3)
sharding=None
] f
i:f32[3,3] = broadcast_in_dim[
broadcast_dimensions=()
shape=(3, 3)
sharding=None
] g
e:f32[3,3] = select_n d i h
in (e,) }
] d 1.0:f32[] 0.0:f32[]
in (e,) }
[[1. 0. 0.]
[0. 1. 0.]
[0. 0. 1.]]
Transpilation to Lean
Our strategy to prove the correctness of this code
is to transpile it over to Lean and then write proofs
for the Lean version of the code. We do this with three
Lean modules.
A mapping of the low-level Jaxpr
An implementation of a Jaxpr interpreter
A mapping of the high-level Jaxpr code.
If we trust our Jaxpr interpreter, we can then
prove things about our JAX code in Lean and then
prove that those properties also apply to the original Python code.
Let's walk through each of the steps.
First, we are going to
map over the Jaxpr from the first section to Lean. Lean has a powerful
macro system, so we can make the mapping
look nearly identical to the original representation. We follow the specified
syntax of the original, propagating forward some additional guarantees about
bounds and dimensions.
Next, we need an implementation of the Jaxpr IR in Lean.
As far as I know, there is not a formal specification of this
language, but each part is well documented. In a more realistic
setting, we would want to match the lower-level instructions exactly.
In addition to trusting that this interpreter matches Jaxpr, we are also
going to relax the floating-point challenges at this stage. Our
interpreter will just pretend things work like the reals.
| .add x y left right _ =>
let xv := x.eval env
let yv := y.eval env
fun i => DType.add _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i))
Given the source of truth of the Jaxpr representation, we can now
transpile the Jaxpr into something that is a bit more readable.
This will look kind of like a line-aligned Lean implementation of the
original code. Since this code will not be executed, we can apply transformations
to make it a bit easier to use for our proofs (for instance, see the variance example
below).
Since this implementation is generated, we shouldn't trust that it is correct.
We therefore need to actually prove this.
The following theorem shows that the code matches its IR.
If we can complete each of these three steps, we know that we have passed the implementation across
the language barrier. Our implementation does this for most of the core Jaxpr IR with the assumption
that we are working with the reals. From here on out, we can assume we are working with implementations
in the proof language Lean.
Puzzle Proofs
We can now go back to the puzzles.
We consider an example of the original outer puzzle, which
asks the user to produce the outer product of two vectors.
Here's the original code and the transpiled Lean.
In a similar way, we convert the loop version to Lean. This one is a bit long since JAX unrolls the loop to a sequence of instructions, so I am omitting the Lean code.
defloop_outer(a, b):
out = jnp.zeros((a.shape[0], b.shape[0]), dtype=a.dtype)for i inrange(a.shape[0]):for j inrange(b.shape[0]):
out = out.at[i, j].set(a[i]* b[j])return out
The goal is to certify that this produces the same output as
the vectorized version. Once we are in Lean, proving this kind of thing
is actually relatively straightforward. The code below automatically converts
the broadcasting operations into expanded cases and enumerates all 2 × 3 cases
to check that they produce the same result as the loop.
We combine this with the proof that the Jaxpr matches the above functions. We now have a full proof that the puzzle was correct!
All goals completed! 🐙
Verified Properties
Of course, just proving the correctness of puzzles is not that useful.
Things become more interesting when we want to complete arbitrary ad hoc
proofs about our code. To illustrate roughly what this might look like,
let's walk through a few properties and look at:
The JAX code
The spec to prove
An example of the proof (that an LLM might write for us)
Transformer Properties
Recently, I wrote a blog post on
Lean-Verified Transformers
that implemented several equivariance properties of transformers directly in Lean.
We probably don't want to write the code in Lean, so we can write it in
JAX and transpile it over.
We first define equivariance as:
B(PX; W) = P B(X; W)
The goal is to prove that a transformer without positional encoding is position-equivariant. We start with the transformer block in JAX.
@jax.jit
deftransformer_block(x, weight, wq, wk, wv):
h = forward(x, weight)
q = project(h, wq)
k = project(h, wk)
v = project(h, wv)return attention(q, k, v)
A proof of part of this property would look like this.
Let's consider another simple, but common, example in ML.
We want to guarantee that, after applying norm clipping in our
code, we arrive at a gradient norm that is at most our fixed bound.
r > 0 ⟹ ‖norm_clip(x, r)‖₂ ≤ min(r, ‖x‖₂)
defnorm_clip(x, radius):# The theorem requires radius > 0, so zero vectors have a safe denominator.return x *(radius / jnp.maximum(radius, jnp.linalg.norm(x)))
The goal here is to obtain the following certificate.
theorem norm_clip_bound (x : Tensor ℝ [3]) (r : ℝ) (hr : 0 < r) :
Tensor.vectorNorm (norm_clip x (Tensor.scalar r)) ≤ min r (Tensor.vectorNorm x)
For a simple example, this is a mechanical calculation. But if you imagine an optimized implementation, we would need to push this through the whole neural network.
theorem l2_clipL2_le (r : ℝ) (hr : 0 < r) (x : ι → ℝ) :
l2 (clipL2 r x) ≤ min r (l2 x) := by
have hd : 0 < max r (l2 x) := lt_of_lt_of_le hr (le_max_left _ _)
have ha : 0 ≤ r / max r (l2 x) := div_nonneg hr.le hd.le
change l2 (fun i => x i * (r / max r (l2 x))) ≤ _
rw [l2_scale, abs_of_nonneg ha]
apply le_min
· calc
_ ≤ (r / max r (l2 x)) * max r (l2 x) :=
mul_le_mul_of_nonneg_left (le_max_right _ _) ha
_ = r := div_mul_cancel₀ _ (ne_of_gt hd)
· calc
_ ≤ 1 * l2 x := mul_le_mul_of_nonneg_right
((div_le_one hd).mpr (le_max_left _ _)) (l2_nonneg x)
_ = _ := one_mul _
Physics
As a final example, we can consider the use of JAX for non-ML projects.
The inspiration for this blog is the
Noether project, which aimed to
guarantee similar properties of mathematical code. The motivating
example is conservation of properties through code.
One thing that I have personally wanted for years is better
guarantees for code that uses sampling, such as RL. Techniques
like probabilistic programming provide these, but are not typically
used in scaled systems. One solution might be to implement the JAX
code as is and prove that it matches a clean spec for its probabilistic
properties.
As a simple example, let us imagine doing Monte Carlo sampling
that takes 16 samples of a die roll. We would like to establish
the variance of our estimator.
As a final example, let's consider how these proofs might actually interact
with all the work going on in formalizing mathematics. This is a toy example
where we verify that the real-valued version of our JAX code produces the mathematically
correct gradient.
We first define the function in JAX and then perform autodiff
in the standard JAX manner to create a function for its derivative.
defquadratic(x):return x * x +3.0* x
defgradient(x):return jax.grad(quadratic)(x)
JAX will then lower the new function to Jaxpr that we can apply our
method to. It goes through the conversion to Lean in the standard way.
{ lambda ; a:f32[]. let
b:f32[] = mul a a
c:f32[] = mul 3.0:f32[] a
_:f32[] = add b c
d:f32[] = mul 3.0:f32[] 1.0:f32[]
e:f32[] = mul a 1.0:f32[]
f:f32[] = add_any d e
g:f32[] = mul 1.0:f32[] a
h:f32[] = add_any f g
in (h,) }
There are many things we might prove here, but to keep things simple
let's just get a certificate that the autodiff Jaxpr is equivalent to the
mathematical derivative of the original function. We are using mathlib
here with the deriv operator to produce that function directly.
Used `tac1 <;> tac2` where `(tac1; tac2)` would sufficeNote: This linter can be disabled with `set_option linter.unnecessarySeqFocus false`theoremquadratic_hasDerivAt(x:ℝ):HasDerivAt(funy:ℝ=>quadratic(Tensor.scalary)())(2*x+3)x:=byx:ℝ⊢ HasDerivAt(funy=>quadratic(Tensor.scalary)())(2*x+3)xhaveh:HasDerivAt(funy:ℝ=>y*y+3*y)(1*x+x*1+3*1)x:=((hasDerivAt_idx).mul(hasDerivAt_idx)).add((hasDerivAt_idx).const_mul3)x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ HasDerivAt(funy=>quadratic(Tensor.scalary)())(2*x+3)xconverthusing1e'_8x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ (funy=>quadratic(Tensor.scalary)())=funy=>y*y+3*ye'_9x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ 2*x+3=1*x+x*1+3*1<;>e'_8x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ (funy=>quadratic(Tensor.scalary)())=funy=>y*y+3*ye'_9x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ 2*x+3=1*x+x*1+3*1simp[quadratic,Tensor.scalar]e'_9x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ 2*x=x+xUsed `tac1 <;> tac2` where `(tac1; tac2)` would sufficeNote: This linter can be disabled with `set_option linter.unnecessarySeqFocus false`<;>e'_9x:ℝh:HasDerivAt(funy=>y*y+3*y)(1*x+x*1+3*1)x⊢ 2*x=x+xringAll goals completed! 🐙
Conclusion
This library provides a working implementation of the conversion between
JAX and Lean with the ability to write non-trivial proofs.
This is meant as a glance at what fully verifiable ML might allow
for checking core properties like faithfulness, numerical stability, and
statistical correctness. However, there is still lots to do here, including
dealing with numerical precision, distributed components, making certificates more readable, and
handling the devex of changing code and proofs. As in many areas, the future here
will be quite unpredictable.