import JaxLean.Stdlib import examples.transformer.generated.Transformer import JaxLean.Stdlib.MatrixRules

Read 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 (scores:Matrix 3 3i:Index [3, 3]⊢ 0 < 3 All goals completed! 🐙) scores iscores:Matrix 3 3row:Fin 3⊢ ∑ col, normalizeRows (fun z => 1 + max z 0) scores (row, col, ()) = 1 exact normalizeRows_sum_one _ score_positive (scores:Matrix 3 3row:Fin 3⊢ 0 < 3 All goals completed! 🐙) scores row

Row-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) := selection:Fin 3 → Fin 3x:Matrix 3 2w:Matrix 2 2⊢ forward (selectRows selection x) w = selectRows selection (forward x w) 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) := perm:Fin 3 ≃ Fin 3scores:Matrix 3 3⊢ normalize (reindexMatrix (⇑perm) (⇑perm) scores) = reindexMatrix (⇑perm) (⇑perm) (normalize scores) 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 := rflAll 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) := 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) 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) := 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) All goals completed! 🐙

This theorem is about the imported Jaxpr itself. Translation certificates connect the function-level proof above to its independent IR evaluator.

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) All goals completed! 🐙endend JaxLean.TransformerJax