import JaxLean.Stdlib
import examples.transformer.generated.Transformer
import JaxLean.Stdlib.MatrixRulesRead this file alongside examples/transformer/code.py.
Each proof crosses one function boundary. Only the normalization specification looks inside scalar arithmetic. Block and network proofs use child theorems. The final theorem transports the property back to the imported Jaxpr.
namespace JaxLean.TransformerJaxopen Tensoropen scoped BigOperatorsnoncomputable sectionabbrev Matrix (n m : Nat) := Tensor ℝ [n, m]
The mathematical meaning of Python's normalize, including its denominator.
theorem normalize_spec (scores : Matrix 3 3) :
normalize scores = normalizeRows (fun z => 1 + max z 0) scores := rflprivate theorem score_positive (z : ℝ) : 0 < 1 + max z 0 := z:ℝ⊢ 0 < 1 + max z 0
z:ℝthis:0 ≤ max z 0⊢ 0 < 1 + max z 0
All goals completed! 🐙The example's weights really are a probability vector in every row.
scores:Matrix 3 3i:Index [3, 3]⊢ 0 ≤ normalizeRows (fun z => 1 + max z 0) scores i
exact normalizeRows_nonneg _ score_positive (by scores:Matrix 3 3i:Index [3, 3]⊢ 0 < 3 decide All goals completed! 🐙) scores i
theorem normalize_sum_one (scores : Matrix 3 3) (row : Fin 3) :
(∑ col, normalize scores (row, col, ())) = 1 := by scores:Matrix 3 3row:Fin 3⊢ ∑ col, normalize scores (row, col, ()) = 1
rw [normalize_spec scores:Matrix 3 3row:Fin 3⊢ ∑ col, normalizeRows (fun z => 1 + max z 0) scores (row, col, ()) = 1 scores:Matrix 3 3row:Fin 3⊢ ∑ col, normalizeRows (fun z => 1 + max z 0) scores (row, col, ()) = 1] scores:Matrix 3 3row:Fin 3⊢ ∑ col, normalizeRows (fun z => 1 + max z 0) scores (row, col, ()) = 1
exact normalizeRows_sum_one _ score_positive (by scores:Matrix 3 3row:Fin 3⊢ 0 < 3 decide All goals completed! 🐙) scores rowRow-local functions permit repeated or dropped rows as well as permutations.
theorem project_select (selection : Fin 3 → Fin 3) (x : Matrix 3 2) (w : Matrix 2 2) :
project (selectRows selection x) w = selectRows selection (project x w) :=
(selectRows_matmul selection x w).symmtheorem forward_select (selection : Fin 3 → Fin 3) (x : Matrix 3 2) (w : Matrix 2 2) :
forward (selectRows selection x) w = selectRows selection (forward x w) := by selection:Fin 3 → Fin 3x:Matrix 3 2w:Matrix 2 2⊢ forward (selectRows selection x) w = selectRows selection (forward x w)
simp only [forward, project_select, selectRows_map] All goals completed! 🐙Attention needs a bijection of token positions, since all keys contribute.
theorem normalize_permute (perm : Fin 3 ≃ Fin 3) (scores : Matrix 3 3) :
normalize (reindexMatrix perm perm scores) =
reindexMatrix perm perm (normalize scores) := by perm:Fin 3 ≃ Fin 3scores:Matrix 3 3⊢ normalize (reindexMatrix (⇑perm) (⇑perm) scores) = reindexMatrix (⇑perm) (⇑perm) (normalize scores)
simp only [normalize_spec, normalizeRows_reindex] All goals completed! 🐙A readable specification at the attention boundary.
theorem attention_spec (q k v : Matrix 3 2) :
attention q k v = matmul (normalize (matmul q (transposeMatrix k))) v := rfl
theorem attention_permute (perm : Fin 3 ≃ Fin 3) (q k v : Matrix 3 2) :
attention (selectRows perm q) (selectRows perm k) (selectRows perm v) =
selectRows perm (attention q k v) := by perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ attention (selectRows (⇑perm) q) (selectRows (⇑perm) k) (selectRows (⇑perm) v) = selectRows (⇑perm) (attention q k v)
rw [attention_spec, perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ (normalize ((selectRows (⇑perm) q).matmul (selectRows (⇑perm) k).transposeMatrix)).matmul (selectRows (⇑perm) v) =
selectRows (⇑perm) (attention q k v) All goals completed! 🐙 attention_spec, perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ (normalize ((selectRows (⇑perm) q).matmul (selectRows (⇑perm) k).transposeMatrix)).matmul (selectRows (⇑perm) v) =
selectRows (⇑perm) ((normalize (matmul q (transposeMatrix k))).matmul v) All goals completed! 🐙 matmul_transpose_select, perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ (normalize (reindexMatrix (⇑perm) (⇑perm) (matmul q (transposeMatrix k)))).matmul (selectRows (⇑perm) v) =
selectRows (⇑perm) ((normalize (matmul q (transposeMatrix k))).matmul v) All goals completed! 🐙
normalize_permute, perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ (reindexMatrix (⇑perm) (⇑perm) (normalize (matmul q (transposeMatrix k)))).matmul (selectRows (⇑perm) v) =
selectRows (⇑perm) ((normalize (matmul q (transposeMatrix k))).matmul v) All goals completed! 🐙 matmul_reindex_contract perm:Fin 3 ≃ Fin 3q:Matrix 3 2k:Matrix 3 2v:Matrix 3 2⊢ selectRows (⇑perm) ((normalize (matmul q (transposeMatrix k))).matmul v) =
selectRows (⇑perm) ((normalize (matmul q (transposeMatrix k))).matmul v) All goals completed! 🐙] All goals completed! 🐙The block proof uses only its children's contracts.
theorem transformer_block_permute (perm : Fin 3 ≃ Fin 3)
(x : Matrix 3 2) (weight wq wk wv : Matrix 2 2) :
transformer_block (selectRows perm x) weight wq wk wv =
selectRows perm (transformer_block x weight wq wk wv) := by perm:Fin 3 ≃ Fin 3x:Matrix 3 2weight:Matrix 2 2wq:Matrix 2 2wk:Matrix 2 2wv:Matrix 2 2⊢ transformer_block (selectRows (⇑perm) x) weight wq wk wv = selectRows (⇑perm) (transformer_block x weight wq wk wv)
simp only [transformer_block, forward_select, project_select, attention_permute] All goals completed! 🐙Two separately parameterized blocks: no attention internals appear here.
theorem transformer_permute (perm : Fin 3 ≃ Fin 3) (x : Matrix 3 2)
(w0 q0 k0 v0 w1 q1 k1 v1 : Matrix 2 2) :
transformer (selectRows perm x) w0 q0 k0 v0 w1 q1 k1 v1 =
selectRows perm (transformer x w0 q0 k0 v0 w1 q1 k1 v1) := by perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 = selectRows (⇑perm) (transformer x w0 q0 k0 v0 w1 q1 k1 v1)
simp only [transformer, transformer_block_permute] All goals completed! 🐙This theorem is about the imported Jaxpr itself. Translation certificates connect the function-level proof above to its independent IR evaluator.
theorem certified_transformer_permute (perm : Fin 3 ≃ Fin 3) (x : Matrix 3 2)
(w0 q0 k0 v0 w1 q1 k1 v1 : Matrix 2 2) :
Jaxpr.Program.eval
(.cons (selectRows perm x) (.cons w0 (.cons q0 (.cons k0 (.cons v0
(.cons w1 (.cons q1 (.cons k1 (.cons v1 .nil))))))))) transformer_ir =
selectRows perm (Jaxpr.Program.eval
(.cons x (.cons w0 (.cons q0 (.cons k0 (.cons v0
(.cons w1 (.cons q1 (.cons k1 (.cons v1 .nil))))))))) transformer_ir) := by perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ Jaxpr.Program.eval
(Jaxpr.Env.cons (selectRows (⇑perm) x)
(Jaxpr.Env.cons w0
(Jaxpr.Env.cons q0
(Jaxpr.Env.cons k0
(Jaxpr.Env.cons v0
(Jaxpr.Env.cons w1 (Jaxpr.Env.cons q1 (Jaxpr.Env.cons k1 (Jaxpr.Env.cons v1 Jaxpr.Env.nil)))))))))
transformer_ir =
selectRows (⇑perm)
(Jaxpr.Program.eval
(Jaxpr.Env.cons x
(Jaxpr.Env.cons w0
(Jaxpr.Env.cons q0
(Jaxpr.Env.cons k0
(Jaxpr.Env.cons v0
(Jaxpr.Env.cons w1 (Jaxpr.Env.cons q1 (Jaxpr.Env.cons k1 (Jaxpr.Env.cons v1 Jaxpr.Env.nil)))))))))
transformer_ir)
rw [transformer_translation_correct, perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 =
selectRows (⇑perm)
(Jaxpr.Program.eval
(Jaxpr.Env.cons x
(Jaxpr.Env.cons w0
(Jaxpr.Env.cons q0
(Jaxpr.Env.cons k0
(Jaxpr.Env.cons v0
(Jaxpr.Env.cons w1 (Jaxpr.Env.cons q1 (Jaxpr.Env.cons k1 (Jaxpr.Env.cons v1 Jaxpr.Env.nil)))))))))
transformer_ir) perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 = selectRows (⇑perm) (transformer x w0 q0 k0 v0 w1 q1 k1 v1) transformer_translation_correct perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 = selectRows (⇑perm) (transformer x w0 q0 k0 v0 w1 q1 k1 v1) perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 = selectRows (⇑perm) (transformer x w0 q0 k0 v0 w1 q1 k1 v1)] perm:Fin 3 ≃ Fin 3x:Matrix 3 2w0:Matrix 2 2q0:Matrix 2 2k0:Matrix 2 2v0:Matrix 2 2w1:Matrix 2 2q1:Matrix 2 2k1:Matrix 2 2v1:Matrix 2 2⊢ transformer (selectRows (⇑perm) x) w0 q0 k0 v0 w1 q1 k1 v1 = selectRows (⇑perm) (transformer x w0 q0 k0 v0 w1 q1 k1 v1)
exact transformer_permute perm x w0 q0 k0 v0 w1 q1 k1 v1 All goals completed! 🐙endend JaxLean.TransformerJax