import JaxLean.Core.RealOps import JaxLean.Core.Indexing

One dtype- and shape-indexed Jaxpr. Floating arithmetic is interpreted over mathematical reals; integer indices use Int32 and predicates use Bool.

namespace JaxLean.Jaxpropen scoped BigOperatorsabbrev Shape := List Natdef coordinate : (s : Shape) → Index s → (axis : Nat) → Fin (s[axis]?.getD 1) | [], _, 0 => ⟨0, x✝:Index []⊢ 0 < [][0]?.getD 1 All goals completed! 🐙⟩ | [], _, _ + 1 => ⟨0, x✝:Index []n✝:ℕ⊢ 0 < [][n✝ + 1]?.getD 1 All goals completed! 🐙⟩ | _ :: _, i, 0 => i.1 | _ :: s, i, n + 1 => coordinate s i.2 ndef broadcastValid : Shape → Shape → List Nat → Prop | [], _, dims => dims = [] | n :: ns, t, axis :: axes => axis < t.length ∧ (n = 1 ∨ n = t[axis]?.getD 1) ∧ broadcastValid ns t axes | _ :: _, _, [] => Falseinstance (s t : Shape) (dims : List Nat) : Decidable (broadcastValid s t dims) := s:Shapet:Shapedims:List ℕ⊢ Decidable (broadcastValid s t dims) induction s generalizing dims with t:Shapedims:List ℕ⊢ Decidable (broadcastValid [] t dims) All goals completed! 🐙 t:Shapen:ℕns:List ℕih:(dims : List ℕ) → Decidable (broadcastValid ns t dims)dims:List ℕ⊢ Decidable (broadcastValid (n :: ns) t dims) cases dims with t:Shapen:ℕns:List ℕih:(dims : List ℕ) → Decidable (broadcastValid ns t dims)⊢ Decidable (broadcastValid (n :: ns) t []) All goals completed! 🐙 t:Shapen:ℕns:List ℕih:(dims : List ℕ) → Decidable (broadcastValid ns t dims)a:ℕaxes:List ℕ⊢ Decidable (broadcastValid (n :: ns) t (a :: axes)) t:Shapen:ℕns:List ℕih:(dims : List ℕ) → Decidable (broadcastValid ns t dims)a:ℕaxes:List ℕ⊢ Decidable (a < List.length t ∧ (n = 1 ∨ n = t[a]?.getD 1) ∧ broadcastValid ns t axes); All goals completed! 🐙def broadcastIndex : (s t : Shape) → (dims : List Nat) → broadcastValid s t dims → Index t → Index s | [], _, _, _, _ => () | n :: ns, t, axis :: axes, h, i => (if hn : n = 1 then ⟨0, n:ℕns:List ℕt:Shapeaxis:ℕaxes:List ℕh:broadcastValid (n :: ns) t (axis :: axes)i:Index thn:n = 1⊢ 0 < n All goals completed! 🐙⟩ else ⟨(coordinate t i axis).val, n:ℕns:List ℕt:Shapeaxis:ℕaxes:List ℕh:broadcastValid (n :: ns) t (axis :: axes)i:Index thn:¬n = 1⊢ ↑(coordinate t i axis) < n n:ℕns:List ℕt:Shapeaxis:ℕaxes:List ℕh:broadcastValid (n :: ns) t (axis :: axes)i:Index thn:¬n = 1this:↑(coordinate t i axis) < t[axis]?.getD 1⊢ ↑(coordinate t i axis) < n; n:ℕns:List ℕt:Shapeaxis:ℕaxes:List ℕh:broadcastValid (n :: ns) t (axis :: axes)i:Index thn:¬n = 1this✝:↑(coordinate t i axis) < t[axis]?.getD 1this:n = t[axis]?.getD 1⊢ ↑(coordinate t i axis) < n; All goals completed! 🐙⟩, broadcastIndex ns t axes h.2.2 i) | _ :: _, _, [], h, _ => False.elim h

The source permutation determines the inverse coordinate map.

abbrev transposeDimensions (s : Shape) (permutation : List Nat) : List Nat := (List.range s.length).map (fun axis => permutation.idxOf axis)

Reverse exactly the axes named by the Jaxpr parameter.

def reverseIndex : (s : Shape) → List Nat → Index s → Nat → Index s | [], _, _, _ => () | _ :: ns, axes, i, axis => (if axis ∈ axes then i.1.rev else i.1, reverseIndex ns axes i.2 (axis + 1))inductive DType where | real | bool | int deriving DecidableEqabbrev DType.denote : DType → Type | .real => ℝ | .bool => Bool | .int => Int32abbrev Ty := DType × List Natabbrev Value (t : Ty) := Tensor t.1.denote t.2inductive Var : List Ty → Ty → Type | here : Var (t :: ctx) t | there : Var ctx t → Var (u :: ctx) tinductive Env : List Ty → Type | nil : Env [] | cons : Value t → Env ctx → Env (t :: ctx)def Env.get : Env ctx → Var ctx t → Value t | .cons x _, .here => x | .cons _ xs, .there v => xs.get vinductive Atom (ctx : List Ty) : Ty → Type | var : Var ctx t → Atom ctx t | tensorLiteral : Value t → Atom ctx t | literal (numerator : Int) (denominator : Nat) (positive : 0 < denominator) : Atom ctx (.real, [])noncomputable def Atom.eval (env : Env ctx) : Atom ctx t → Value t | .var v => env.get v | .tensorLiteral x => x | .literal n d _ => fun _ => (n : ℝ) / (d : ℝ)inductive Comparison where | eq | ne | lt | le | gt | gedef Comparison.eval [LinearOrder α] (c : Comparison) (a b : α) : Bool := match c with | .eq => decide (a = b) | .ne => decide (a ≠ b) | .lt => decide (a < b) | .le => decide (a ≤ b) | .gt => decide (a > b) | .ge => decide (a ≥ b)

Comparisons on signed machine indices use their signed mathematical values.

def Comparison.intEval (c : Comparison) (a b : Int32) : Bool := c.eval a.toInt b.toIntnoncomputable def Comparison.scalar (c : Comparison) : (d : DType) → d.denote → d.denote → Bool | .real => c.eval | .int => c.intEval | .bool => fun a b => c.eval a.toNat b.toNatinductive Conversion : DType → DType → Type | identity : Conversion d d | intToReal : Conversion .int .realnoncomputable def Conversion.eval : Conversion a b → a.denote → b.denote | .identity => id | .intToReal => fun x => (x.toInt : ℝ)noncomputable def DType.add : (d : DType) → d.denote → d.denote → d.denote | .real => fun a b => a + b | .int => fun a b => a + b | .bool => fun _ _ => false -- excluded by the primitive's numeric witnessnoncomputable def DType.sub : (d : DType) → d.denote → d.denote → d.denote | .real => fun a b => a - b | .int => fun a b => a - b | .bool => fun _ _ => false -- excluded by the primitive's numeric witnessnoncomputable def DType.mul : (d : DType) → d.denote → d.denote → d.denote | .real => fun a b => a * b | .int => fun a b => a * b | .bool => fun _ _ => false -- excluded by the primitive's numeric witnessinductive Args (ctx : List Ty) : List Ty → Type | nil : Args ctx [] | cons : Atom ctx t → Args ctx ts → Args ctx (t :: ts)noncomputable def Args.eval (env : Env ctx) : Args ctx ts → Env ts | .nil => .nil | .cons x xs => .cons (x.eval env) (xs.eval env)

Read a variadic operand at the selected bounded coordinate.

def Env.read (env : Env types) : ((s : Shape) × Var types (d, s) × Index s) → d.denote | ⟨_, v, j⟩ => let value := env.get v value jtheorem Env.read_dite (env : Env types) (p : Prop) [Decidable p] (yes : p → (s : Shape) × Var types (d, s) × Index s) (no : ¬p → (s : Shape) × Var types (d, s) × Index s) : env.read (if h : p then yes h else no h) = if h : p then env.read (yes h) else env.read (no h) := types:List Tyd:DTypeenv:Env typesp:Propinst✝:Decidable pyes:p → (s : Shape) × Var types (d, s) × Index sno:¬p → (s : Shape) × Var types (d, s) × Index s⊢ env.read (if h : p then yes h else no h) = if h : p then env.read (yes h) else env.read (no h) types:List Tyd:DTypeenv:Env typesp:Propinst✝:Decidable pyes:p → (s : Shape) × Var types (d, s) × Index sno:¬p → (s : Shape) × Var types (d, s) × Index sh✝:p⊢ env.read (yes h✝) = env.read (yes h✝)types:List Tyd:DTypeenv:Env typesp:Propinst✝:Decidable pyes:p → (s : Shape) × Var types (d, s) × Index sno:¬p → (s : Shape) × Var types (d, s) × Index sh✝:¬p⊢ env.read (no h✝) = env.read (no h✝) types:List Tyd:DTypeenv:Env typesp:Propinst✝:Decidable pyes:p → (s : Shape) × Var types (d, s) × Index sno:¬p → (s : Shape) × Var types (d, s) × Index sh✝:p⊢ env.read (yes h✝) = env.read (yes h✝)types:List Tyd:DTypeenv:Env typesp:Propinst✝:Decidable pyes:p → (s : Shape) × Var types (d, s) × Index sno:¬p → (s : Shape) × Var types (d, s) × Index sh✝:¬p⊢ env.read (no h✝) = env.read (no h✝) All goals completed! 🐙

Jaxpr primitives. Layout maps are bounded; scatter plans are validated static metadata.

inductive Op (ctx : List Ty) : Ty → Type | iota (shape : Shape) (dimension : Nat) (valid : dimension < shape.length := by decide) : Op ctx (.real, shape) | exp : Atom ctx (.real, s) → Op ctx (.real, s) | log : Atom ctx (.real, s) → Op ctx (.real, s) | sqrt : Atom ctx (.real, s) → Op ctx (.real, s) | rsqrt : Atom ctx (.real, s) → Op ctx (.real, s) | sin : Atom ctx (.real, s) → Op ctx (.real, s) | cos : Atom ctx (.real, s) → Op ctx (.real, s) | tanh : Atom ctx (.real, s) → Op ctx (.real, s) | dot_general (left : Index t → Index k → Index s) (right : Index t → Index k → Index u) : Atom ctx (.real, s) → Atom ctx (.real, u) → Op ctx (.real, t) | reduce_sum (map : Index t → Index k → Index s) : Atom ctx (.real, s) → Op ctx (.real, t) | reduce_max (positive : 0 < n) (map : Index t → Fin n → Index s) : Atom ctx (.real, s) → Op ctx (.real, t) | reduce_min (positive : 0 < n) (map : Index t → Fin n → Index s) : Atom ctx (.real, s) → Op ctx (.real, t) | concatenate (args : Args ctx types) (select : Index t → (s : Shape) × Var types (d, s) × Index s) : Op ctx (d, t) | copy : Atom ctx (d, s) → Op ctx (d, s) | stop_gradient : Atom ctx (d, s) → Op ctx (d, s) | add (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) (numeric : d ≠ .bool := by decide) : Op ctx (d, t) | sub (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) (numeric : d ≠ .bool := by decide) : Op ctx (d, t) | neg : Atom ctx (.real, s) → Op ctx (.real, s) | square : Atom ctx (.real, s) → Op ctx (.real, s) | integer_pow (power : Int) : Atom ctx (.real, s) → Op ctx (.real, s) | abs : Atom ctx (.real, s) → Op ctx (.real, s) | min (x : Atom ctx (.real, s)) (y : Atom ctx (.real, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.real, t) | max (x : Atom ctx (.real, s)) (y : Atom ctx (.real, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.real, t) | mul (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) (numeric : d ≠ .bool := by decide) : Op ctx (d, t) | div (x : Atom ctx (.real, s)) (y : Atom ctx (.real, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.real, t) | broadcast_in_dim (shape : Shape) (broadcast_dimensions : List Nat) (x : Atom ctx (d, s)) (valid : broadcastValid s shape broadcast_dimensions := by decide) : Op ctx (d, shape) | transpose (permutation : List Nat) (x : Atom ctx (d, s)) (valid : t = permutation.map (fun axis => s[axis]?.getD 1) ∧ permutation.Perm (List.range s.length) ∧ broadcastValid s t (transposeDimensions s permutation) := by decide) : Op ctx (d, t) | squeeze (map : Index t → Index s) : Atom ctx (d, s) → Op ctx (d, t) | slice (map : Index t → Index s) : Atom ctx (d, s) → Op ctx (d, t) | rev (dimensions : List Nat) (x : Atom ctx (d, s)) (valid : dimensions.Nodup ∧ dimensions.all (· < s.length) := by decide) : Op ctx (d, s) | reshape (map : Index t → Index s) : Atom ctx (d, s) → Op ctx (d, t) | scatter_add (plan : List (Index s × Index u)) : Atom ctx (.real, s) → Atom ctx (.real, u) → Op ctx (.real, s) | scatter (plan : List (Index s × Index u)) : Atom ctx (.real, s) → Atom ctx (.real, u) → Op ctx (.real, s) | eq (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | ne (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | lt (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | le (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | gt (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | ge (x : Atom ctx (d, s)) (y : Atom ctx (d, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | and (x : Atom ctx (.bool, s)) (y : Atom ctx (.bool, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | or (x : Atom ctx (.bool, s)) (y : Atom ctx (.bool, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | xor (x : Atom ctx (.bool, s)) (y : Atom ctx (.bool, u)) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (.bool, t) | not : Atom ctx (.bool, s) → Op ctx (.bool, s) | select_n (condition : Atom ctx (.bool, c)) (no : Atom ctx (d, s)) (yes : Atom ctx (d, u)) (predicate : broadcastValid c t (List.range c.length) := by decide) (left : broadcastValid s t (List.range s.length) := by decide) (right : broadcastValid u t (List.range u.length) := by decide) : Op ctx (d, t) | gather (map : Tensor Int32 u → Index t → Index s) : Atom ctx (.real, s) → Atom ctx (.int, u) → Op ctx (.real, t) | convert_element_type (new_dtype : DType) (x : Atom ctx (d, s)) (conversion : Conversion d new_dtype := by first | exact .identity | exact .intToReal) : Op ctx (new_dtype, s)noncomputable def Op.eval (env : Env ctx) : Op ctx t → Value t | .iota shape dimension _ => fun i => ((coordinate shape i dimension).val : ℝ) | .exp x => let xv := x.eval env fun i => Real.exp (xv i) | .log x => let xv := x.eval env fun i => Real.log (xv i) | .sqrt x => let xv := x.eval env fun i => Real.sqrt (xv i) | .rsqrt x => let xv := x.eval env fun i => (Real.sqrt (xv i))⁻¹ | .sin x => let xv := x.eval env fun i => Real.sin (xv i) | .cos x => let xv := x.eval env fun i => Real.cos (xv i) | .tanh x => let xv := x.eval env fun i => (Real.exp (xv i) - Real.exp (-xv i)) / (Real.exp (xv i) + Real.exp (-xv i)) | .dot_general left right x y => let yv := y.eval env let xv := x.eval env fun i => ∑ j, xv (left i j) * yv (right i j) | .reduce_sum map x => let xv := x.eval env fun i => ∑ j, xv (map i j) | .reduce_max positive map x => let xv := x.eval env fun i => (List.ofFn (fun j => xv (map i j))).foldl Max.max (xv (map i ⟨0, positive⟩)) | .reduce_min positive map x => let xv := x.eval env fun i => (List.ofFn (fun j => xv (map i j))).foldl Min.min (xv (map i ⟨0, positive⟩)) | .concatenate args select => fun i => (args.eval env).read (select i) | .stop_gradient x => let xv := x.eval env xv | .copy x => let xv := x.eval env xv | .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)) | .sub x y left right _ => let xv := x.eval env let yv := y.eval env fun i => DType.sub _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .neg x => let xv := x.eval env fun i => -(xv i) | .square x => let xv := x.eval env fun i => xv i ^ (2 : Nat) | .integer_pow p x => let xv := x.eval env fun i => xv i ^ p | .abs x => let xv := x.eval env fun i => |xv i| | .min x y left right => let yv := y.eval env let xv := x.eval env fun i => Min.min (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .max x y left right => let yv := y.eval env let xv := x.eval env fun i => Max.max (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .mul x y left right _ => let xv := x.eval env let yv := y.eval env fun i => DType.mul _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .div x y left right => let yv := y.eval env let xv := x.eval env fun i => xv (broadcastIndex _ _ _ left i) / yv (broadcastIndex _ _ _ right i) | @Op.broadcast_in_dim _ d source shape dims x h => let xv := x.eval env fun i => xv (broadcastIndex source shape dims h i) | @Op.transpose _ d source target _permutation x valid => fun i => (Atom.eval (t := (d, source)) env x) (broadcastIndex source target _ valid.2.2 i) | @Op.squeeze _ target source d map x => fun i => (Atom.eval (t := (d, source)) env x) (map i) | @Op.slice _ target source d map x => fun i => (Atom.eval (t := (d, source)) env x) (map i) | .rev dimensions x _ => fun i => x.eval env (reverseIndex _ dimensions i 0) | @Op.reshape _ target source d map x => fun i => (Atom.eval (t := (d, source)) env x) (map i) | .scatter_add plan x update => let xv := x.eval env let uv := update.eval env plan.foldl (fun tensor (position, source) => fun i => if i = position then tensor i + uv source else tensor i) xv | .scatter plan x update => let xv := x.eval env let uv := update.eval env plan.foldl (fun tensor (position, source) => fun i => if i = position then uv source else tensor i) xv | .eq x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .eq _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .ne x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .ne _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .lt x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .lt _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .le x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .le _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .gt x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .gt _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .ge x y left right => let xv := x.eval env let yv := y.eval env fun i => Comparison.scalar .ge _ (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .and x y left right => let yv := y.eval env let xv := x.eval env fun i => (xv (broadcastIndex _ _ _ left i)) && (yv (broadcastIndex _ _ _ right i)) | .or x y left right => let yv := y.eval env let xv := x.eval env fun i => (xv (broadcastIndex _ _ _ left i)) || (yv (broadcastIndex _ _ _ right i)) | .xor x y left right => let yv := y.eval env let xv := x.eval env fun i => Bool.xor (xv (broadcastIndex _ _ _ left i)) (yv (broadcastIndex _ _ _ right i)) | .not x => let xv := x.eval env fun i => !(xv i) | .select_n c no yes predicate left right => let cv := c.eval env let nv := no.eval env let yv := yes.eval env fun i => if cv (broadcastIndex _ _ _ predicate i) then yv (broadcastIndex _ _ _ right i) else nv (broadcastIndex _ _ _ left i) | @Op.gather _ _u target source map x indices => fun i => (Atom.eval (t := (.real, source)) env x) (map (indices.eval env) i) | .convert_element_type _ x conversion => let xv := x.eval env fun i => conversion.eval (xv i)inductive Program : List Ty → Ty → Type | ret : Atom ctx t → Program ctx t | bind : Op ctx t → Program (t :: ctx) u → Program ctx u | call : Program args t → Args ctx args → Program (t :: ctx) u → Program ctx unoncomputable def Program.eval (env : Env ctx) : Program ctx t → Value t | .ret x => x.eval env | .bind op rest => let value := op.eval env let next := Env.cons value env rest.eval next | .call body args rest => let value := body.eval (args.eval env) let next := Env.cons value env rest.eval nextend JaxLean.Jaxpr