Compare commits
13 Commits
a5f533d67a
...
86bc33ee26
| Author | SHA1 | Date | |
|---|---|---|---|
| 86bc33ee26 | |||
| 9e0702b5f5 | |||
| 445187837c | |||
| 1a49689edc | |||
| b1b3b0d2fe | |||
| 379438ec17 | |||
| 1120e01605 | |||
| b6b30958aa | |||
| 5737805125 | |||
| e738eb4294 | |||
| 6a6ed521ca | |||
| c38c10fe9e | |||
| c367f130cf |
@@ -19,3 +19,10 @@ import Spa.Showable
|
||||
import Spa.Analysis.Utils
|
||||
import Spa.Analysis.Sign
|
||||
import Spa.Analysis.Constant
|
||||
import Spa.Language.Tagged.Id
|
||||
import Spa.Language.Tagged.Derive
|
||||
import Spa.Language.Tagged.Basic
|
||||
import Spa.Language.Tagged.Properties
|
||||
import Spa.Language.Tagged.Graphs
|
||||
import Spa.Analysis.Reaching
|
||||
import Spa.Transformation.Licm
|
||||
|
||||
@@ -29,15 +29,13 @@ def minus : ConstLattice → ConstLattice → ConstLattice
|
||||
|
||||
lemma plus_mono₂ : Monotone₂ plus :=
|
||||
AboveBelow.monotone₂_of_strict plus
|
||||
(fun y => by cases y <;> rfl) (fun x => by cases x <;> rfl)
|
||||
(fun y hy => by cases y <;> first | exact absurd rfl hy | rfl)
|
||||
(fun x hx => by cases x <;> first | exact absurd rfl hx | rfl)
|
||||
(fun y => by aesop) (fun x => by aesop)
|
||||
(fun y hy => by aesop) (fun x hx => by aesop)
|
||||
|
||||
lemma minus_mono₂ : Monotone₂ minus :=
|
||||
AboveBelow.monotone₂_of_strict minus
|
||||
(fun y => by cases y <;> rfl) (fun x => by cases x <;> rfl)
|
||||
(fun y hy => by cases y <;> first | exact absurd rfl hy | rfl)
|
||||
(fun x hx => by cases x <;> first | exact absurd rfl hx | rfl)
|
||||
(fun y => by aesop) (fun x => by aesop)
|
||||
(fun y hy => by aesop) (fun x hx => by aesop)
|
||||
|
||||
def interpConst : ConstLattice → Value → Prop
|
||||
| .bot, _ => False
|
||||
@@ -96,36 +94,14 @@ def output : String :=
|
||||
lemma plus_valid {g₁ g₂ : ConstLattice} {z₁ z₂ : ℤ}
|
||||
(h₁ : ⟦g₁⟧ (.int z₁)) (h₂ : ⟦g₂⟧ (.int z₂)) :
|
||||
⟦plus g₁ g₂⟧ (.int (z₁ + z₂)) := by
|
||||
rcases g₁ with _ | _ | c₁
|
||||
· exact h₁.elim
|
||||
· rcases g₂ with _ | _ | c₂
|
||||
· exact h₂.elim
|
||||
· exact trivial
|
||||
· exact trivial
|
||||
· rcases g₂ with _ | _ | c₂
|
||||
· exact h₂.elim
|
||||
· exact trivial
|
||||
· injection h₁ with hz₁
|
||||
injection h₂ with hz₂
|
||||
show Value.int (z₁ + z₂) = Value.int (c₁ + c₂)
|
||||
rw [hz₁, hz₂]
|
||||
rcases g₁ with _ | _ | c₁ <;> rcases g₂ with _ | _ | c₂ <;>
|
||||
simp_all [plus, constInterpretation, interpConst]
|
||||
|
||||
lemma minus_valid {g₁ g₂ : ConstLattice} {z₁ z₂ : ℤ}
|
||||
(h₁ : ⟦g₁⟧ (.int z₁)) (h₂ : ⟦g₂⟧ (.int z₂)) :
|
||||
⟦minus g₁ g₂⟧ (.int (z₁ - z₂)) := by
|
||||
rcases g₁ with _ | _ | c₁
|
||||
· exact h₁.elim
|
||||
· rcases g₂ with _ | _ | c₂
|
||||
· exact h₂.elim
|
||||
· exact trivial
|
||||
· exact trivial
|
||||
· rcases g₂ with _ | _ | c₂
|
||||
· exact h₂.elim
|
||||
· exact trivial
|
||||
· injection h₁ with hz₁
|
||||
injection h₂ with hz₂
|
||||
show Value.int (z₁ - z₂) = Value.int (c₁ - c₂)
|
||||
rw [hz₁, hz₂]
|
||||
rcases g₁ with _ | _ | c₁ <;> rcases g₂ with _ | _ | c₂ <;>
|
||||
simp_all [minus, constInterpretation, interpConst]
|
||||
|
||||
instance eval_valid : ValidExprEvaluator ConstLattice prog := by
|
||||
constructor
|
||||
@@ -158,7 +134,7 @@ instance eval_valid : ValidExprEvaluator ConstLattice prog := by
|
||||
exact minus_valid h₁ h₂
|
||||
|
||||
theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result ConstLattice prog) ⟧ ρ :=
|
||||
⟦ variablesAt prog.finalState (result ConstLattice prog) ⟧ ρ () :=
|
||||
Forward.analyze_correct ConstLattice prog hrun
|
||||
|
||||
end ConstAnalysis
|
||||
|
||||
@@ -9,13 +9,22 @@ namespace Forward
|
||||
|
||||
variable {L : Type} [FiniteHeightLattice L] {prog : Program} [E : StmtEvaluator L prog]
|
||||
|
||||
def evalStmtOrNone (s : prog.State) (o : Option BasicStmt) (hco : prog.code s = o)
|
||||
(vs : VariableValues L prog) : VariableValues L prog :=
|
||||
o.elimEq vs (fun bs h => E.eval s bs (hco.trans h))
|
||||
|
||||
lemma evalStmtOrNone_mono (s : prog.State) (o : Option BasicStmt)
|
||||
(hco : prog.code s = o) : Monotone (evalStmtOrNone (L := L) s o hco) :=
|
||||
elimEq_self_mono o (fun bs h vs => E.eval s bs (hco.trans h) vs)
|
||||
(fun bs h => E.eval_mono s bs (hco.trans h))
|
||||
|
||||
def updateVariablesForState (s : prog.State) (sv : StateVariables L prog) :
|
||||
VariableValues L prog :=
|
||||
(prog.code s).foldl (fun vs bs => E.eval s bs vs) (variablesAt s sv)
|
||||
evalStmtOrNone s (prog.code s) rfl (variablesAt s sv)
|
||||
|
||||
lemma updateVariablesForState_mono (s : prog.State) :
|
||||
Monotone (updateVariablesForState (L := L) s) := fun _ _ hle =>
|
||||
foldl_mono' (prog.code s) _ (E.eval_mono s ·) (variablesAt_le hle s)
|
||||
evalStmtOrNone_mono s (prog.code s) rfl (variablesAt_le hle s)
|
||||
|
||||
def updateAll (sv : StateVariables L prog) : StateVariables L prog :=
|
||||
FiniteMap.generalizedUpdate id updateVariablesForState
|
||||
@@ -54,67 +63,99 @@ lemma joinForKey_initialState :
|
||||
rw [joinForKey, prog.incoming_initialState_eq_nil]
|
||||
rfl
|
||||
|
||||
variable [I : LatticeInterpretation L] [V : ValidStmtEvaluator L prog]
|
||||
class ValidStateEvaluator (L : Type) [FiniteHeightLattice L] (prog : Program)
|
||||
[E : StmtEvaluator L prog] [S : StateInterp L prog] where
|
||||
step : (s : prog.State) → {ρ₁ ρ₂ : Env} → {bs : BasicStmt} →
|
||||
prog.code s = some bs → EvalBasicStmt ρ₁ bs ρ₂ → S.St ρ₁ → S.St ρ₂
|
||||
valid : ∀ (s : prog.State) {ρ₁ ρ₂ : Env} {bs : BasicStmt}
|
||||
{vs : VariableValues L prog} {st : S.St ρ₁},
|
||||
(hcode : prog.code s = some bs) → (hbs : EvalBasicStmt ρ₁ bs ρ₂) → ⟦ vs ⟧ ρ₁ st →
|
||||
⟦ E.eval s bs hcode vs ⟧ ρ₂ (step s hcode hbs st)
|
||||
botV_init : ⟦ botV L prog ⟧ [] S.init
|
||||
|
||||
instance [LatticeInterpretation L] [ValidStmtEvaluator L prog] :
|
||||
ValidStateEvaluator L prog where
|
||||
step := by intro _ _ _ _ _ _ _; exact PUnit.unit
|
||||
valid := by intro _ _ _ _ _ _ hcode hbs hvs; exact ValidStmtEvaluator.valid hcode hbs hvs
|
||||
botV_init := by intro k l _ v hmem; cases hmem
|
||||
|
||||
section
|
||||
variable [S : StateInterp L prog] [V : ValidStateEvaluator L prog]
|
||||
|
||||
noncomputable def stepStmtOrNone (s : prog.State) {ρ₁ ρ₂ : Env} :
|
||||
(o : Option BasicStmt) → prog.code s = o → EvalBasicStmtOpt ρ₁ o ρ₂ →
|
||||
S.St ρ₁ → S.St ρ₂
|
||||
| none, _, .none, st => st
|
||||
| some _, hco, .some hbs, st => V.step s hco hbs st
|
||||
|
||||
noncomputable def stepNode (s : prog.State) {ρ₁ ρ₂ : Env}
|
||||
(h : EvalBasicStmtOpt ρ₁ (prog.code s) ρ₂) (st : S.St ρ₁) : S.St ρ₂ :=
|
||||
stepStmtOrNone s (prog.code s) rfl h st
|
||||
|
||||
noncomputable def stepTraceState :
|
||||
{s₁ s₂ : prog.State} → {ρ₁ ρ₂ : Env} →
|
||||
Trace prog.cfg s₁ s₂ ρ₁ ρ₂ → S.St ρ₁ → S.St ρ₂
|
||||
| s₁, _, _, _, .single hnode, st => stepNode s₁ hnode st
|
||||
| s₁, _, _, _, .edge hnode _ subtr, st =>
|
||||
stepTraceState subtr (stepNode s₁ hnode st)
|
||||
|
||||
omit [DecidableEq L] in
|
||||
lemma eval_fold_valid {s : prog.State} {bss : List BasicStmt}
|
||||
{vs : VariableValues L prog} {ρ₁ ρ₂ : Env}
|
||||
(hbss : EvalBasicStmts ρ₁ bss ρ₂) (hvs : ⟦ vs ⟧ ρ₁) :
|
||||
⟦ bss.foldl (fun vs bs => E.eval s bs vs) vs ⟧ ρ₂ := by
|
||||
induction hbss generalizing vs with
|
||||
| nil => exact hvs
|
||||
| cons hbs _ ih => exact ih (ValidStmtEvaluator.valid hbs hvs)
|
||||
|
||||
omit [DecidableEq L] in
|
||||
lemma updateVariablesForState_matches {s : prog.State}
|
||||
{sv : StateVariables L prog} {ρ₁ ρ₂ : Env}
|
||||
(hbss : EvalBasicStmts ρ₁ (prog.code s) ρ₂)
|
||||
(hvs : ⟦ variablesAt s sv ⟧ ρ₁) :
|
||||
⟦ updateVariablesForState s sv ⟧ ρ₂ :=
|
||||
eval_fold_valid hbss hvs
|
||||
lemma evalStmtOrNone_valid {s : prog.State} {ρ₁ ρ₂ : Env} {st : S.St ρ₁}
|
||||
{vs : VariableValues L prog} (o : Option BasicStmt) (hco : prog.code s = o)
|
||||
(he : EvalBasicStmtOpt ρ₁ o ρ₂) (hvs : ⟦ vs ⟧ ρ₁ st) :
|
||||
⟦ evalStmtOrNone s o hco vs ⟧ ρ₂ (stepStmtOrNone s o hco he st) := by
|
||||
cases he with
|
||||
| none => exact hvs
|
||||
| some hbs => exact V.valid s hco hbs hvs
|
||||
|
||||
omit [DecidableEq L] in
|
||||
lemma updateAll_matches {s : prog.State} {sv : StateVariables L prog}
|
||||
{ρ₁ ρ₂ : Env} (hbss : EvalBasicStmts ρ₁ (prog.code s) ρ₂)
|
||||
(hvs : ⟦ variablesAt s sv ⟧ ρ₁) :
|
||||
⟦ variablesAt s (updateAll sv) ⟧ ρ₂ := by
|
||||
{ρ₁ ρ₂ : Env} {st : S.St ρ₁}
|
||||
(hnode : EvalBasicStmtOpt ρ₁ (prog.code s) ρ₂)
|
||||
(hvs : ⟦ variablesAt s sv ⟧ ρ₁ st) :
|
||||
⟦ variablesAt s (updateAll sv) ⟧ ρ₂ (stepNode s hnode st) := by
|
||||
rw [variablesAt_updateAll]
|
||||
exact updateVariablesForState_matches hbss hvs
|
||||
exact evalStmtOrNone_valid (prog.code s) rfl hnode hvs
|
||||
|
||||
lemma stepTrace {s₁ : prog.State} {ρ₁ ρ₂ : Env}
|
||||
(hjoin : ⟦ joinForKey s₁ (result L prog) ⟧ ρ₁)
|
||||
(hbss : EvalBasicStmts ρ₁ (prog.code s₁) ρ₂) :
|
||||
⟦ variablesAt s₁ (result L prog) ⟧ ρ₂ := by
|
||||
lemma stepTrace {s₁ : prog.State} {ρ₁ ρ₂ : Env} {st : S.St ρ₁}
|
||||
(hjoin : ⟦ joinForKey s₁ (result L prog) ⟧ ρ₁ st)
|
||||
(hnode : EvalBasicStmtOpt ρ₁ (prog.code s₁) ρ₂) :
|
||||
⟦ variablesAt s₁ (result L prog) ⟧ ρ₂ (stepNode s₁ hnode st) := by
|
||||
rw [result_eq L prog]
|
||||
refine updateAll_matches hbss ?_
|
||||
refine updateAll_matches hnode ?_
|
||||
rw [variablesAt_joinAll]
|
||||
exact hjoin
|
||||
|
||||
lemma walkTrace {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env}
|
||||
(hjoin : ⟦ joinForKey s₁ (result L prog) ⟧ ρ₁)
|
||||
lemma walkTrace {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env} {st₁ : S.St ρ₁}
|
||||
(hjoin : ⟦ joinForKey s₁ (result L prog) ⟧ ρ₁ st₁)
|
||||
(tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂) :
|
||||
⟦ variablesAt s₂ (result L prog) ⟧ ρ₂ := by
|
||||
⟦ variablesAt s₂ (result L prog) ⟧ ρ₂ (stepTraceState tr st₁) := by
|
||||
induction tr with
|
||||
| single hbss => exact stepTrace hjoin hbss
|
||||
| @edge _ ρ' _ i₁ i₂ _ hbss hedge _ ih =>
|
||||
have hstep : ⟦ variablesAt i₁ (result L prog) ⟧ ρ' :=
|
||||
stepTrace hjoin hbss
|
||||
| single hnode => exact stepTrace hjoin hnode
|
||||
| @edge _ ρ' _ i₁ i₂ _ hnode hedge _ ih =>
|
||||
have hstep : ⟦ variablesAt i₁ (result L prog) ⟧ ρ' (stepNode i₁ hnode st₁) :=
|
||||
stepTrace hjoin hnode
|
||||
have hmem : variablesAt i₁ (result L prog)
|
||||
∈ (result L prog).valuesAt (prog.incoming i₂) :=
|
||||
FiniteMap.mem_valuesAt prog.states_nodup
|
||||
(prog.mem_incoming_of_edge hedge) (variablesAt_mem i₁ (result L prog))
|
||||
exact ih (interp_foldr hstep hmem)
|
||||
|
||||
omit V in
|
||||
lemma interp_joinForKey_initialState :
|
||||
⟦ joinForKey prog.initialState (result L prog) ⟧ [] := by
|
||||
variable (L prog) in
|
||||
theorem analyze_correct_state {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result L prog) ⟧ ρ
|
||||
(stepTraceState (prog.trace hrun) S.init) := by
|
||||
refine walkTrace ?_ (prog.trace hrun)
|
||||
rw [joinForKey_initialState]
|
||||
exact interp_botV_nil
|
||||
exact ValidStateEvaluator.botV_init
|
||||
|
||||
end
|
||||
|
||||
variable (L prog) in
|
||||
theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result L prog) ⟧ ρ :=
|
||||
walkTrace interp_joinForKey_initialState (prog.trace hrun)
|
||||
theorem analyze_correct [LatticeInterpretation L] [ValidStmtEvaluator L prog]
|
||||
{ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result L prog) ⟧ ρ () :=
|
||||
analyze_correct_state L prog hrun
|
||||
|
||||
end Forward
|
||||
|
||||
|
||||
@@ -14,14 +14,14 @@ lemma updateVariablesFromExpression_mono (k : String) (e : Expr) :
|
||||
Monotone (updateVariablesFromExpression (L := L) (prog := prog) k e) :=
|
||||
FiniteMap.generalizedUpdate_monotone monotone_id (fun _ => E.eval_mono e)
|
||||
|
||||
def evalBasicStmt (_ : prog.State) (bs : BasicStmt)
|
||||
def evalBasicStmt (s : prog.State) (bs : BasicStmt) (_h : prog.code s = some bs)
|
||||
(vs : VariableValues L prog) : VariableValues L prog :=
|
||||
match bs with
|
||||
| .assign k e => updateVariablesFromExpression k e vs
|
||||
| .noop => vs
|
||||
|
||||
lemma evalBasicStmt_mono (s : prog.State) (bs : BasicStmt) :
|
||||
Monotone (evalBasicStmt (L := L) (prog := prog) s bs) := by
|
||||
lemma evalBasicStmt_mono (s : prog.State) (bs : BasicStmt) (h : prog.code s = some bs) :
|
||||
Monotone (evalBasicStmt (L := L) (prog := prog) s bs h) := by
|
||||
cases bs with
|
||||
| assign k e => exact updateVariablesFromExpression_mono k e
|
||||
| noop => exact monotone_id
|
||||
@@ -32,7 +32,7 @@ instance ExprEvaluator.toStmtEvaluator : StmtEvaluator L prog :=
|
||||
instance ExprEvaluator.toStmtEvaluator_valid [LatticeInterpretation L]
|
||||
[ValidExprEvaluator L prog] : ValidStmtEvaluator L prog := by
|
||||
constructor
|
||||
intro s vs ρ₁ ρ₂ bs hbs hvs
|
||||
intro s vs ρ₁ ρ₂ bs hcode hbs hvs
|
||||
cases hbs with
|
||||
| noop => exact hvs
|
||||
| assign k e v hev =>
|
||||
|
||||
@@ -7,8 +7,9 @@ namespace Forward
|
||||
variable (L : Type) [Lattice L] (prog : Program)
|
||||
|
||||
class StmtEvaluator where
|
||||
eval : prog.State → BasicStmt → VariableValues L prog → VariableValues L prog
|
||||
eval_mono : ∀ s bs, Monotone (eval s bs)
|
||||
eval : (s : prog.State) → (bs : BasicStmt) → prog.code s = some bs →
|
||||
VariableValues L prog → VariableValues L prog
|
||||
eval_mono : ∀ s bs h, Monotone (eval s bs h)
|
||||
|
||||
class ExprEvaluator where
|
||||
eval : Expr → VariableValues L prog → L
|
||||
@@ -17,13 +18,13 @@ class ExprEvaluator where
|
||||
class ValidExprEvaluator [ExprEvaluator L prog] [I : LatticeInterpretation L] :
|
||||
Prop where
|
||||
valid : ∀ {vs : VariableValues L prog} {ρ : Env} {e : Expr} {v : Value},
|
||||
EvalExpr ρ e v → ⟦ vs ⟧ ρ → I.interp (ExprEvaluator.eval e vs) v
|
||||
EvalExpr ρ e v → ⟦ vs ⟧ ρ () → I.interp (ExprEvaluator.eval e vs) v
|
||||
|
||||
class ValidStmtEvaluator [E : StmtEvaluator L prog] [LatticeInterpretation L] :
|
||||
Prop where
|
||||
valid : ∀ {s : prog.State} {vs : VariableValues L prog} {ρ₁ ρ₂ : Env}
|
||||
{bs : BasicStmt},
|
||||
EvalBasicStmt ρ₁ bs ρ₂ → ⟦ vs ⟧ ρ₁ → ⟦ E.eval s bs vs ⟧ ρ₂
|
||||
{bs : BasicStmt} (hcode : prog.code s = some bs),
|
||||
EvalBasicStmt ρ₁ bs ρ₂ → ⟦ vs ⟧ ρ₁ () → ⟦ E.eval s bs hcode vs ⟧ ρ₂ ()
|
||||
|
||||
end Forward
|
||||
|
||||
|
||||
@@ -64,39 +64,47 @@ lemma variablesAt_joinAll (s : prog.State) (sv : StateVariables L prog) :
|
||||
variablesAt s (joinAll sv) = joinForKey s sv :=
|
||||
joinAll_mem_eq (variablesAt_mem s (joinAll sv))
|
||||
|
||||
/-! ### Lifting an interpretation to variable maps -/
|
||||
class StateInterp (L : Type) [Lattice L] (prog : Program) where
|
||||
St : Env → Type
|
||||
init : St []
|
||||
interp : VariableValues L prog → (ρ : Env) → St ρ → Prop
|
||||
interp_sup : ∀ {vs₁ vs₂ : VariableValues L prog} {ρ : Env} {st : St ρ},
|
||||
interp vs₁ ρ st ∨ interp vs₂ ρ st → interp (vs₁ ⊔ vs₂) ρ st
|
||||
interp_inf : ∀ {vs₁ vs₂ : VariableValues L prog} {ρ : Env} {st : St ρ},
|
||||
interp vs₁ ρ st ∧ interp vs₂ ρ st → interp (vs₁ ⊓ vs₂) ρ st
|
||||
|
||||
variable [I : LatticeInterpretation L]
|
||||
instance [S : StateInterp L prog] :
|
||||
Interp (VariableValues L prog) ((ρ : Env) → S.St ρ → Prop) :=
|
||||
⟨S.interp⟩
|
||||
|
||||
omit [FiniteHeightLattice L] in
|
||||
instance : Interp (VariableValues L prog) (Env → Prop) where
|
||||
interp (vs : VariableValues L prog) (ρ : Env) : Prop :=
|
||||
∀ (k : String) (l : L), (k, l) ∈ vs →
|
||||
∀ (v : Value), Env.Mem (k, v) ρ → I.interp l v
|
||||
|
||||
lemma interp_botV_nil : ⟦ botV L prog ⟧ [] := by
|
||||
intro k l _ v hmem
|
||||
cases hmem
|
||||
|
||||
omit [FiniteHeightLattice L] in
|
||||
lemma interp_sup {vs₁ vs₂ : VariableValues L prog} {ρ : Env}
|
||||
(h : ⟦ vs₁⟧ ρ ∨ ⟦ vs₂ ⟧ ρ) : ⟦ vs₁ ⊔ vs₂ ⟧ ρ := by
|
||||
intro k l hmem v hv
|
||||
obtain ⟨l₁, l₂, rfl, h₁, h₂⟩ := FiniteMap.mem_sup hmem
|
||||
rcases h with h | h
|
||||
· exact I.interp_sup v (Or.inl (h _ _ h₁ _ hv))
|
||||
· exact I.interp_sup v (Or.inr (h _ _ h₂ _ hv))
|
||||
|
||||
lemma interp_foldr {vs : VariableValues L prog}
|
||||
{vss : List (VariableValues L prog)} {ρ : Env}
|
||||
(hvs : ⟦ vs ⟧ ρ) (hmem : vs ∈ vss) :
|
||||
⟦ vss.foldr (· ⊔ ·) (botV L prog) ⟧ ρ := by
|
||||
lemma interp_foldr [S : StateInterp L prog]
|
||||
{vs : VariableValues L prog} {vss : List (VariableValues L prog)}
|
||||
{ρ : Env} {st : S.St ρ} (hvs : ⟦ vs ⟧ ρ st) (hmem : vs ∈ vss) :
|
||||
⟦ vss.foldr (· ⊔ ·) (botV L prog) ⟧ ρ st := by
|
||||
induction vss with
|
||||
| nil => cases hmem
|
||||
| cons vs' vss' ih =>
|
||||
rcases List.mem_cons.mp hmem with rfl | hmem'
|
||||
· exact interp_sup (Or.inl hvs)
|
||||
· exact interp_sup (Or.inr (ih hmem'))
|
||||
· exact S.interp_sup (Or.inl hvs)
|
||||
· exact S.interp_sup (Or.inr (ih hmem'))
|
||||
|
||||
variable [I : LatticeInterpretation L]
|
||||
|
||||
instance : StateInterp L prog where
|
||||
St := fun _ => PUnit
|
||||
init := PUnit.unit
|
||||
interp vs ρ _ := ∀ (k : String) (l : L), (k, l) ∈ vs →
|
||||
∀ (v : Value), Env.Mem (k, v) ρ → I.interp l v
|
||||
interp_sup := by
|
||||
intro vs₁ vs₂ ρ st h k l hmem v hv
|
||||
obtain ⟨l₁, l₂, rfl, h₁, h₂⟩ := FiniteMap.mem_sup hmem
|
||||
rcases h with h | h
|
||||
· exact I.interp_sup v (Or.inl (h _ _ h₁ _ hv))
|
||||
· exact I.interp_sup v (Or.inr (h _ _ h₂ _ hv))
|
||||
interp_inf := by
|
||||
intro vs₁ vs₂ ρ st h k l hmem v hv
|
||||
obtain ⟨l₁, l₂, rfl, h₁, h₂⟩ := FiniteMap.mem_inf hmem
|
||||
exact I.interp_inf v ⟨h.1 _ _ h₁ _ hv, h.2 _ _ h₂ _ hv⟩
|
||||
|
||||
end Forward
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import Spa.Analysis.Forward
|
||||
import Spa.Lattice.Bool
|
||||
import Spa.Lattice.Tuple
|
||||
import Spa.Language.Tagged.Graphs
|
||||
import Spa.Showable
|
||||
|
||||
namespace Spa
|
||||
@@ -8,23 +10,31 @@ open Forward
|
||||
|
||||
instance : Showable Bool := ⟨fun b => if b then "true" else "false"⟩
|
||||
|
||||
abbrev DefSet (prog : Program) : Type := FiniteMap prog.State Bool prog.states
|
||||
instance {n : ℕ} {β : Type*} [Showable β] : Showable (Fin n → β) :=
|
||||
⟨fun f =>
|
||||
"{" ++ (List.finRange n).foldr
|
||||
(fun i rest => show' i ++ " ↦ " ++ show' (f i) ++ ", " ++ rest) ""
|
||||
++ "}"⟩
|
||||
|
||||
abbrev DefSet (prog : Program) : Type := prog.NodeId → Bool
|
||||
|
||||
namespace ReachingAnalysis
|
||||
|
||||
variable (prog : Program)
|
||||
|
||||
def genSet (s : prog.State) : DefSet prog :=
|
||||
FiniteMap.updating (⊥ : DefSet prog) [s] (fun _ => true)
|
||||
def genSet (s : prog.State) {bs : BasicStmt} (h : prog.code s = some bs) :
|
||||
DefSet prog :=
|
||||
Function.update (⊥ : DefSet prog) (prog.nodeIdOfNonempty s h) true
|
||||
|
||||
def eval (s : prog.State) :
|
||||
BasicStmt → VariableValues (DefSet prog) prog → VariableValues (DefSet prog) prog
|
||||
| .assign k _, vs =>
|
||||
FiniteMap.generalizedUpdate id (fun _ _ => genSet prog s) [k] vs
|
||||
| .noop, vs => vs
|
||||
(bs : BasicStmt) → prog.code s = some bs →
|
||||
VariableValues (DefSet prog) prog → VariableValues (DefSet prog) prog
|
||||
| .assign k _, h, vs =>
|
||||
FiniteMap.generalizedUpdate id (fun _ _ => genSet prog s h) [k] vs
|
||||
| .noop, _, vs => vs
|
||||
|
||||
lemma eval_mono (s : prog.State) (bs : BasicStmt) :
|
||||
Monotone (eval prog s bs) := by
|
||||
lemma eval_mono (s : prog.State) (bs : BasicStmt) (h : prog.code s = some bs) :
|
||||
Monotone (eval prog s bs h) := by
|
||||
cases bs with
|
||||
| assign k e =>
|
||||
exact FiniteMap.generalizedUpdate_monotone monotone_id (fun _ => monotone_const)
|
||||
@@ -36,6 +46,59 @@ instance stmtEvaluator : StmtEvaluator (DefSet prog) prog :=
|
||||
def output : String :=
|
||||
show' (result (DefSet prog) prog)
|
||||
|
||||
inductive Run (prog : Program) where
|
||||
| nil : Run prog
|
||||
| cons (s : prog.State) (bs : BasicStmt) (hc : prog.code s = some bs)
|
||||
(rest : Run prog) : Run prog
|
||||
|
||||
@[aesop unsafe cases]
|
||||
inductive LastAssign (prog : Program) (x : String) : Run prog → prog.NodeId → Prop
|
||||
| here (s : prog.State) (e : Expr) (hc : prog.code s = some (.assign x e))
|
||||
(rest : Run prog) :
|
||||
LastAssign prog x (Run.cons s (.assign x e) hc rest) (prog.nodeIdOfNonempty s hc)
|
||||
| there (s : prog.State) (bs : BasicStmt) (hc : prog.code s = some bs)
|
||||
(rest : Run prog) {n : prog.NodeId} :
|
||||
(∀ e, bs ≠ .assign x e) → LastAssign prog x rest n →
|
||||
LastAssign prog x (Run.cons s bs hc rest) n
|
||||
|
||||
instance stateInterp : StateInterp (DefSet prog) prog where
|
||||
St := fun _ => Run prog
|
||||
init := Run.nil
|
||||
interp vs _ run := ∀ (x : String) (assigners : DefSet prog), (x, assigners) ∈ vs →
|
||||
∀ (n : prog.NodeId), LastAssign prog x run n → assigners n = true
|
||||
interp_sup := by
|
||||
intro vs₁ vs₂ ρ run h x assigners hmem n hla
|
||||
obtain ⟨a₁, a₂, rfl, h₁, h₂⟩ := FiniteMap.mem_sup hmem
|
||||
aesop
|
||||
interp_inf := by
|
||||
intro vs₁ vs₂ ρ run h x assigners hmem n hla
|
||||
obtain ⟨a₁, a₂, rfl, h₁, h₂⟩ := FiniteMap.mem_inf hmem
|
||||
aesop
|
||||
|
||||
instance validStateEvaluator : ValidStateEvaluator (DefSet prog) prog where
|
||||
step := by intro s _ _ bs hcode _ rest; exact Run.cons s bs hcode rest
|
||||
valid := by
|
||||
intro s ρ₁ ρ₂ bs vs st hcode hbs hvs
|
||||
cases hbs with
|
||||
| noop => intro x assigners hmem n hla; aesop
|
||||
| assign x e v hev =>
|
||||
intro k assigners hmem n hla
|
||||
have hmem2 : (k, assigners) ∈
|
||||
FiniteMap.generalizedUpdate id (fun _ _ => genSet prog s hcode) [x] vs := hmem
|
||||
by_cases hx : k = x
|
||||
· subst hx
|
||||
have hd := FiniteMap.generalizedUpdate_mem_eq (List.mem_singleton.mpr rfl) hmem2
|
||||
aesop (add simp genSet)
|
||||
· have hmem' := FiniteMap.generalizedUpdate_not_mem_backward
|
||||
(fun hc => hx (List.mem_singleton.mp hc)) hmem2
|
||||
aesop
|
||||
botV_init := by intro x assigners _ n hla; cases hla
|
||||
|
||||
theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result (DefSet prog) prog) ⟧ ρ
|
||||
(stepTraceState (prog.trace hrun) (stateInterp prog).init) :=
|
||||
Forward.analyze_correct_state (DefSet prog) prog hrun
|
||||
|
||||
end ReachingAnalysis
|
||||
|
||||
end Spa
|
||||
|
||||
@@ -13,6 +13,8 @@ inductive Sign where
|
||||
| zero
|
||||
deriving DecidableEq
|
||||
|
||||
attribute [aesop safe cases] Sign
|
||||
|
||||
instance : Showable Sign :=
|
||||
⟨fun
|
||||
| .plus => "+"
|
||||
@@ -57,21 +59,13 @@ def minus : SignLattice → SignLattice → SignLattice
|
||||
|
||||
lemma plus_mono₂ : Monotone₂ plus :=
|
||||
AboveBelow.monotone₂_of_strict plus
|
||||
(fun y => by cases y <;> rfl)
|
||||
(fun x => by rcases x with _ | _ | s <;> first | rfl | (cases s <;> rfl))
|
||||
(fun y hy => by cases y <;> first | exact absurd rfl hy | rfl)
|
||||
(fun x hx => by
|
||||
rcases x with _ | _ | s <;>
|
||||
first | exact absurd rfl hx | rfl | (cases s <;> rfl))
|
||||
(fun y => by aesop) (fun x => by aesop)
|
||||
(fun y hy => by aesop) (fun x hx => by aesop)
|
||||
|
||||
lemma minus_mono₂ : Monotone₂ minus :=
|
||||
AboveBelow.monotone₂_of_strict minus
|
||||
(fun y => by cases y <;> rfl)
|
||||
(fun x => by rcases x with _ | _ | s <;> first | rfl | (cases s <;> rfl))
|
||||
(fun y hy => by cases y <;> first | exact absurd rfl hy | rfl)
|
||||
(fun x hx => by
|
||||
rcases x with _ | _ | s <;>
|
||||
first | exact absurd rfl hx | rfl | (cases s <;> rfl))
|
||||
(fun y => by aesop) (fun x => by aesop)
|
||||
(fun y hy => by aesop) (fun x hx => by aesop)
|
||||
|
||||
def interpSign : SignLattice → Value → Prop
|
||||
| .bot, _ => False
|
||||
@@ -216,7 +210,7 @@ instance eval_valid : ValidExprEvaluator SignLattice prog := by
|
||||
exact minus_valid h₁ h₂
|
||||
|
||||
theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
|
||||
⟦ variablesAt prog.finalState (result SignLattice prog) ⟧ ρ :=
|
||||
⟦ variablesAt prog.finalState (result SignLattice prog) ⟧ ρ () :=
|
||||
Forward.analyze_correct SignLattice prog hrun
|
||||
|
||||
end SignAnalysis
|
||||
|
||||
@@ -12,7 +12,7 @@ def doStep (f : α → α) (hf : Monotone f) :
|
||||
∀ (g : ℕ) (c : LTSeries α), c.length + g = height (α := α) + 1 →
|
||||
c.last ≤ f c.last → {a : α // a = f a}
|
||||
| 0, c, hlen, _ =>
|
||||
absurd (FiniteHeightLattice.chains_bounded c) (by simp only [height] at hlen; omega)
|
||||
absurd (FiniteHeightLattice.chains_bounded c) (by omega)
|
||||
| g + 1, c, hlen, hle =>
|
||||
if heq : c.last = f c.last then
|
||||
⟨c.last, heq⟩
|
||||
@@ -24,8 +24,7 @@ def doStep (f : α → α) (hf : Monotone f) :
|
||||
def fix (f : α → α) (hf : Monotone f) : {a : α // a = f a} :=
|
||||
doStep f hf (height (α := α) + 1) (RelSeries.singleton _ ⊥)
|
||||
(by simp)
|
||||
(by simpa [RelSeries.last_singleton]
|
||||
using FiniteHeightLattice.bot_le α (f ⊥))
|
||||
(by simp)
|
||||
|
||||
def aFix (f : α → α) (hf : Monotone f) : α :=
|
||||
(fix f hf).1
|
||||
@@ -40,7 +39,7 @@ lemma doStep_le (f : α → α) (hf : Monotone f)
|
||||
(hle : c.last ≤ f c.last), c.last ≤ b →
|
||||
(doStep f hf g c hlen hle : α) ≤ b
|
||||
| 0, c, hlen, _ => fun _ =>
|
||||
absurd (FiniteHeightLattice.chains_bounded c) (by simp only [height] at hlen; omega)
|
||||
absurd (FiniteHeightLattice.chains_bounded c) (by omega)
|
||||
| g + 1, c, hlen, hle => fun hcb => by
|
||||
rw [doStep]
|
||||
split
|
||||
@@ -50,7 +49,7 @@ lemma doStep_le (f : α → α) (hf : Monotone f)
|
||||
|
||||
theorem aFix_le (f : α → α) (hf : Monotone f)
|
||||
{a : α} (ha : a = f a) : aFix f hf ≤ a :=
|
||||
doStep_le f hf ha _ _ _ _ (by simpa using FiniteHeightLattice.bot_le α a)
|
||||
doStep_le f hf ha _ _ _ _ (by simp)
|
||||
|
||||
end Fixedpoint
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ def initialState : p.State := p.rootStmt.cfg.wrapInput
|
||||
|
||||
def finalState : p.State := p.rootStmt.cfg.wrapOutput
|
||||
|
||||
theorem trace {ρ : Env} (h : EvalStmt [] p.rootStmt ρ) :
|
||||
noncomputable def trace {ρ : Env} (h : EvalStmt [] p.rootStmt ρ) :
|
||||
Trace p.cfg p.initialState p.finalState [] ρ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := EndToEndTrace.wrap (Stmt.cfg_sufficient h)
|
||||
rw [Graph.wrap_inputs, List.mem_singleton] at h₁
|
||||
@@ -41,7 +41,7 @@ lemma states_complete (s : p.State) : s ∈ p.states := p.cfg.mem_indices s
|
||||
|
||||
lemma states_nodup : p.states.Nodup := p.cfg.nodup_indices
|
||||
|
||||
def code (st : p.State) : List BasicStmt := p.cfg.nodes st
|
||||
def code (st : p.State) : Option BasicStmt := p.cfg.nodes st
|
||||
|
||||
def incoming (s : p.State) : List p.State := p.cfg.predecessors s
|
||||
|
||||
|
||||
@@ -130,9 +130,9 @@ def loopOut (g : GGraph α) : Fin (2 + g.size) := (1 : Fin 2).castAdd g.size
|
||||
|
||||
This is technically sloppy (see module comment), but it's simple.
|
||||
-/
|
||||
def loop (g : GGraph (List β)) : GGraph (List β) where
|
||||
def loop (g : GGraph (Option β)) : GGraph (Option β) where
|
||||
size := 2 + g.size
|
||||
nodes := Fin.append (fun _ : Fin 2 => []) g.nodes
|
||||
nodes := Fin.append (fun _ : Fin 2 => none) g.nodes
|
||||
edges := g.edges.finNatAddProd 2 ++
|
||||
((g.loopIn, ·) <$> g.inputs.finNatAdd 2) ++
|
||||
((·, g.loopOut) <$> g.outputs.finNatAdd 2) ++
|
||||
@@ -140,9 +140,9 @@ def loop (g : GGraph (List β)) : GGraph (List β) where
|
||||
inputs := [g.loopIn]
|
||||
outputs := [g.loopOut]
|
||||
|
||||
@[simp] lemma loop_inputs (g : GGraph (List β)) : (loop g).inputs = [g.loopIn] := rfl
|
||||
@[simp] lemma loop_inputs (g : GGraph (Option β)) : (loop g).inputs = [g.loopIn] := rfl
|
||||
|
||||
@[simp] lemma loop_outputs (g : GGraph (List β)) : (loop g).outputs = [g.loopOut] := rfl
|
||||
@[simp] lemma loop_outputs (g : GGraph (Option β)) : (loop g).outputs = [g.loopOut] := rfl
|
||||
|
||||
/-- Creates a single-node graph whose node contains the given value. -/
|
||||
def singleton (a : α) : GGraph α where
|
||||
@@ -154,8 +154,8 @@ def singleton (a : α) : GGraph α where
|
||||
|
||||
/-- Creates a new graph with a single input and single output node. Useful to ensure there's
|
||||
a single point of entry and single point of exit. -/
|
||||
def wrap (g : GGraph (List β)) : GGraph (List β) :=
|
||||
singleton [] ⤳ g ⤳ singleton []
|
||||
def wrap (g : GGraph (Option β)) : GGraph (Option β) :=
|
||||
singleton none ⤳ g ⤳ singleton none
|
||||
|
||||
@[simp] lemma map_singleton (f : α → β) (a : α) :
|
||||
f <$> singleton a = singleton (f a) := rfl
|
||||
@@ -176,16 +176,16 @@ def wrap (g : GGraph (List β)) : GGraph (List β) :=
|
||||
funext i
|
||||
refine Fin.addCases ?_ ?_ i <;> intro j <;> simp [Fin.append_left, Fin.append_right]
|
||||
|
||||
@[simp] lemma map_loop (h : β → γ) (g : GGraph (List β)) :
|
||||
(List.map h) <$> (loop g) = loop (List.map h <$> g) := by
|
||||
@[simp] lemma map_loop (h : β → γ) (g : GGraph (Option β)) :
|
||||
(Option.map h) <$> (loop g) = loop (Option.map h <$> g) := by
|
||||
rcases g with ⟨n, nd, e, i, o⟩
|
||||
simp only [Functor.map, GGraph.loop]
|
||||
congr 1
|
||||
funext i
|
||||
refine Fin.addCases ?_ ?_ i <;> intro j <;> simp [Fin.append_left, Fin.append_right]
|
||||
|
||||
@[simp] lemma map_wrap (h : β → γ) (g : GGraph (List β)) :
|
||||
(List.map h) <$> wrap g = wrap (List.map h <$> g) := by
|
||||
@[simp] lemma map_wrap (h : β → γ) (g : GGraph (Option β)) :
|
||||
(Option.map h) <$> wrap g = wrap (Option.map h <$> g) := by
|
||||
simp [GGraph.wrap, GGraph.map_sequence, GGraph.map_singleton]
|
||||
|
||||
variable (g : GGraph α)
|
||||
@@ -220,8 +220,8 @@ lemma edge_of_mem_predecessors {idx₁ idx₂ : g.Index}
|
||||
end GGraph
|
||||
|
||||
/-- "Normal" graphs, for the purposes of the analyses in this
|
||||
framework, have basic blocks in their nodes, and nothing else. -/
|
||||
abbrev Graph : Type := GGraph (List BasicStmt)
|
||||
framework, have basic statements in their nodes, and nothing else. -/
|
||||
abbrev Graph : Type := GGraph (Option BasicStmt)
|
||||
|
||||
namespace Graph
|
||||
|
||||
@@ -235,7 +235,7 @@ end Graph
|
||||
open Graph in
|
||||
def Stmt.cfg : Stmt → Graph
|
||||
-- A basic statement goes into a single basic block
|
||||
| .basic bs => singleton [bs]
|
||||
| .basic bs => singleton (some bs)
|
||||
-- Sequencing of statements corresponds naturally to CFG sequencing
|
||||
| .andThen s₁ s₂ => s₁.cfg ⤳ s₂.cfg
|
||||
-- An if can execute either one branch or the other; overlap them.
|
||||
|
||||
@@ -1,5 +1,24 @@
|
||||
import Spa.Language.Traces
|
||||
|
||||
/-!
|
||||
|
||||
# Properties of the Object Language, CFGs, and Traces
|
||||
|
||||
This module encodes some properties of the language, mostly those having to do
|
||||
with connecting the computational view (the `Spa.Graph`s, on which static
|
||||
analyses are executed) to the semantic view (such as `EvalStmt`, which
|
||||
encodes the expected formal behavior of the language). In particular,
|
||||
to prove that our computationally-implemented static analyses are correct,
|
||||
we need to show that our computational model of their execution (the CFG)
|
||||
matches the formal description. Thus, the key result `cfg_sufficient`.
|
||||
|
||||
Many lemmas and definitions here aim are used to prove that result,
|
||||
by allowing inductive proofs on the construction of the CFG:
|
||||
the bits where we _build up_ the trace corresponding to each
|
||||
proof tree are exactly those when we have two graphs (through
|
||||
which traces exist) and we want to combine these graphs, while
|
||||
showing also that a combined trace exists as well. -/
|
||||
|
||||
namespace Spa
|
||||
|
||||
open Graph
|
||||
@@ -11,13 +30,13 @@ lemma Fin.castAdd_ne_natAdd {n m : ℕ} (i : Fin n) (j : Fin m) :
|
||||
simp only [Fin.coe_castAdd, Fin.coe_natAdd] at this
|
||||
omega
|
||||
|
||||
/-! ### Trace embeddings -/
|
||||
|
||||
section Embeddings
|
||||
|
||||
variable {g₁ g₂ : Graph} {ρ₁ ρ₂ : Env}
|
||||
|
||||
lemma Trace.overlay_left {idx₁ idx₂ : g₁.Index}
|
||||
/-- When two graphs are overlaid, for each trace in the left graph,
|
||||
a corresponding trace exists in the combined graph. -/
|
||||
noncomputable def Trace.overlay_left {idx₁ idx₂ : g₁.Index}
|
||||
(tr : Trace g₁ idx₁ idx₂ ρ₁ ρ₂) :
|
||||
Trace (g₁ ∙ g₂) (idx₁.castAdd g₂.size) (idx₂.castAdd g₂.size) ρ₁ ρ₂ := by
|
||||
induction tr with
|
||||
@@ -29,7 +48,9 @@ lemma Trace.overlay_left {idx₁ idx₂ : g₁.Index}
|
||||
· rwa [show (g₁ ∙ g₂).nodes = Fin.append g₁.nodes g₂.nodes from rfl, Fin.append_left]
|
||||
· exact List.mem_append_left _ (List.mem_map_of_mem _ he)
|
||||
|
||||
lemma Trace.overlay_right {idx₁ idx₂ : g₂.Index}
|
||||
/-- When two graphs are overlaid, for each trace in the right graph,
|
||||
a corresponding trace exists in the combined graph. -/
|
||||
noncomputable def Trace.overlay_right {idx₁ idx₂ : g₂.Index}
|
||||
(tr : Trace g₂ idx₁ idx₂ ρ₁ ρ₂) :
|
||||
Trace (g₁ ∙ g₂) (idx₁.natAdd g₁.size) (idx₂.natAdd g₁.size) ρ₁ ρ₂ := by
|
||||
induction tr with
|
||||
@@ -41,7 +62,9 @@ lemma Trace.overlay_right {idx₁ idx₂ : g₂.Index}
|
||||
· rwa [show (g₁ ∙ g₂).nodes = Fin.append g₁.nodes g₂.nodes from rfl, Fin.append_right]
|
||||
· exact List.mem_append_right _ (List.mem_map_of_mem _ he)
|
||||
|
||||
lemma Trace.sequence_left {idx₁ idx₂ : g₁.Index}
|
||||
/-- When two graphs are sequenced, for each trace in the first graph,
|
||||
a corresponding trace exists in the combined graph. -/
|
||||
noncomputable def Trace.sequence_left {idx₁ idx₂ : g₁.Index}
|
||||
(tr : Trace g₁ idx₁ idx₂ ρ₁ ρ₂) :
|
||||
Trace (g₁ ⤳ g₂) (idx₁.castAdd g₂.size) (idx₂.castAdd g₂.size) ρ₁ ρ₂ := by
|
||||
induction tr with
|
||||
@@ -53,7 +76,9 @@ lemma Trace.sequence_left {idx₁ idx₂ : g₁.Index}
|
||||
· rwa [show (g₁ ⤳ g₂).nodes = Fin.append g₁.nodes g₂.nodes from rfl, Fin.append_left]
|
||||
· exact List.mem_append_left _ (List.mem_append_left _ (List.mem_map_of_mem _ he))
|
||||
|
||||
lemma Trace.sequence_right {idx₁ idx₂ : g₂.Index}
|
||||
/-- When two graphs are sequenced, for each trace in the second graph,
|
||||
a corresponding trace exists in the combined graph. -/
|
||||
noncomputable def Trace.sequence_right {idx₁ idx₂ : g₂.Index}
|
||||
(tr : Trace g₂ idx₁ idx₂ ρ₁ ρ₂) :
|
||||
Trace (g₁ ⤳ g₂) (idx₁.natAdd g₁.size) (idx₂.natAdd g₁.size) ρ₁ ρ₂ := by
|
||||
induction tr with
|
||||
@@ -66,61 +91,72 @@ lemma Trace.sequence_right {idx₁ idx₂ : g₂.Index}
|
||||
· exact List.mem_append_left _
|
||||
(List.mem_append_right _ (List.mem_map_of_mem _ he))
|
||||
|
||||
lemma EndToEndTrace.overlay_left (etr : EndToEndTrace g₁ ρ₁ ρ₂) :
|
||||
/-- Equivalent of `Trace.overlay_left` for end-to-end traces. -/
|
||||
noncomputable def EndToEndTrace.overlay_left (etr : EndToEndTrace g₁ ρ₁ ρ₂) :
|
||||
EndToEndTrace (g₁ ∙ g₂) ρ₁ ρ₂ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := etr
|
||||
exact ⟨i₁.castAdd g₂.size, List.mem_append_left _ (List.mem_map_of_mem _ h₁),
|
||||
i₂.castAdd g₂.size, List.mem_append_left _ (List.mem_map_of_mem _ h₂),
|
||||
tr.overlay_left⟩
|
||||
|
||||
lemma EndToEndTrace.overlay_right (etr : EndToEndTrace g₂ ρ₁ ρ₂) :
|
||||
/-- Equivalent of `Trace.overlay_right` for end-to-end traces. -/
|
||||
noncomputable def EndToEndTrace.overlay_right (etr : EndToEndTrace g₂ ρ₁ ρ₂) :
|
||||
EndToEndTrace (g₁ ∙ g₂) ρ₁ ρ₂ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := etr
|
||||
exact ⟨i₁.natAdd g₁.size, List.mem_append_right _ (List.mem_map_of_mem _ h₁),
|
||||
i₂.natAdd g₁.size, List.mem_append_right _ (List.mem_map_of_mem _ h₂),
|
||||
tr.overlay_right⟩
|
||||
|
||||
lemma EndToEndTrace.concat {ρ₃ : Env} (etr₁ : EndToEndTrace g₁ ρ₁ ρ₂)
|
||||
/-- When two graphs are sequenced, two end-to-end traces through the respective
|
||||
graphs can be sequenced to create an end-to-end trace in the combined
|
||||
graph. This is only possible for end-to-end traces and not for general
|
||||
`Trace`s, because sequencing only introduces edges from the output nodes
|
||||
of one graph to the input nodes of another graph. A non-end-to-end trace
|
||||
need to conclude at the output node, so it cannot necessarily be sequenced
|
||||
with a trace in another graph. -/
|
||||
noncomputable def EndToEndTrace.concat {ρ₃ : Env} (etr₁ : EndToEndTrace g₁ ρ₁ ρ₂)
|
||||
(etr₂ : EndToEndTrace g₂ ρ₂ ρ₃) : EndToEndTrace (g₁ ⤳ g₂) ρ₁ ρ₃ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁
|
||||
obtain ⟨j₁, k₁, j₂, k₂, tr₂⟩ := etr₂
|
||||
refine ⟨i₁.castAdd g₂.size, List.mem_map_of_mem _ h₁,
|
||||
j₂.natAdd g₁.size, List.mem_map_of_mem _ k₂,
|
||||
Trace.concat tr₁.sequence_left ?_ tr₂.sequence_right⟩
|
||||
tr₁.sequence_left ++< ?_ >++ tr₂.sequence_right⟩
|
||||
exact List.mem_append_right _
|
||||
(List.mem_product.mpr ⟨List.mem_map_of_mem _ h₂, List.mem_map_of_mem _ k₁⟩)
|
||||
|
||||
end Embeddings
|
||||
|
||||
/-! ### Loops -/
|
||||
|
||||
section Loop
|
||||
|
||||
variable {g : Graph} {ρ₁ ρ₂ ρ₃ : Env}
|
||||
|
||||
lemma Trace.loop {idx₁ idx₂ : g.Index} (tr : Trace g idx₁ idx₂ ρ₁ ρ₂) :
|
||||
/-- A trace through a body CFG still exists (up to reindexing) in a zero-or-more loop CFG. -/
|
||||
noncomputable def Trace.loop {idx₁ idx₂ : g.Index} (tr : Trace g idx₁ idx₂ ρ₁ ρ₂) :
|
||||
Trace (Graph.loop g) (idx₁.natAdd 2) (idx₂.natAdd 2) ρ₁ ρ₂ := by
|
||||
induction tr with
|
||||
| single hbs =>
|
||||
exact Trace.single (by
|
||||
rwa [show (Graph.loop g).nodes = Fin.append (fun _ : Fin 2 => []) g.nodes from rfl,
|
||||
rwa [show (Graph.loop g).nodes = Fin.append (fun _ : Fin 2 => none) g.nodes from rfl,
|
||||
Fin.append_right])
|
||||
| edge hbs he _ ih =>
|
||||
refine Trace.edge ?_ ?_ ih
|
||||
· rwa [show (Graph.loop g).nodes = Fin.append (fun _ : Fin 2 => []) g.nodes from rfl,
|
||||
· rwa [show (Graph.loop g).nodes = Fin.append (fun _ : Fin 2 => none) g.nodes from rfl,
|
||||
Fin.append_right]
|
||||
· exact List.mem_append_left _ (List.mem_append_left _
|
||||
(List.mem_append_left _ (List.mem_map_of_mem _ he)))
|
||||
|
||||
/-- The beginning node of a loop graph is empty. -/
|
||||
private lemma loop_nodes_at_in :
|
||||
(Graph.loop g).nodes g.loopIn = [] :=
|
||||
Fin.append_left (fun _ : Fin 2 => []) g.nodes 0
|
||||
(Graph.loop g).nodes g.loopIn = none :=
|
||||
Fin.append_left (fun _ : Fin 2 => none) g.nodes 0
|
||||
|
||||
/-- The ending node of a loop graph is empty. -/
|
||||
private lemma loop_nodes_at_out :
|
||||
(Graph.loop g).nodes g.loopOut = [] :=
|
||||
Fin.append_left (fun _ : Fin 2 => []) g.nodes 1
|
||||
(Graph.loop g).nodes g.loopOut = none :=
|
||||
Fin.append_left (fun _ : Fin 2 => none) g.nodes 1
|
||||
|
||||
lemma EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) :
|
||||
/-- Equivlaent of `Trace.loop` for end-to-end traces. -/
|
||||
noncomputable def EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) :
|
||||
EndToEndTrace (Graph.loop g) ρ₁ ρ₂ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := etr
|
||||
-- the edge in → (2 ↑ʳ i₁), reached through the second edge group
|
||||
@@ -132,15 +168,17 @@ lemma EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) :
|
||||
refine List.mem_append_left _ (List.mem_append_right _ ?_)
|
||||
exact List.mem_map_of_mem _ (List.mem_map_of_mem _ h₂)
|
||||
refine ⟨g.loopIn, List.mem_singleton_self _, g.loopOut, List.mem_singleton_self _, ?_⟩
|
||||
exact Trace.concat (Trace.single (loop_nodes_at_in ▸ EvalBasicStmts.nil)) hin
|
||||
(Trace.concat tr.loop hout (Trace.single (loop_nodes_at_out ▸ EvalBasicStmts.nil)))
|
||||
exact Trace.single (loop_nodes_at_in ▸ EvalBasicStmtOpt.none) ++< hin >++
|
||||
tr.loop ++< hout >++ Trace.single (loop_nodes_at_out ▸ EvalBasicStmtOpt.none)
|
||||
|
||||
/-- The zero-or-more times loop has an edge to return back to the top, to continue after an iteration. -/
|
||||
private lemma loop_edge_out_in :
|
||||
((g.loopOut, g.loopIn) : (Graph.loop g).Edge) ∈ (Graph.loop g).edges := by
|
||||
refine List.mem_append_right _ ?_
|
||||
exact List.mem_cons_self _ _
|
||||
|
||||
lemma EndToEndTrace.loop_concat (etr₁ : EndToEndTrace (Graph.loop g) ρ₁ ρ₂)
|
||||
/-- Two traces through a loop can be combined, since a loop can be executed any number of times. -/
|
||||
noncomputable def EndToEndTrace.loop_concat (etr₁ : EndToEndTrace (Graph.loop g) ρ₁ ρ₂)
|
||||
(etr₂ : EndToEndTrace (Graph.loop g) ρ₂ ρ₃) :
|
||||
EndToEndTrace (Graph.loop g) ρ₁ ρ₃ := by
|
||||
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁
|
||||
@@ -148,37 +186,55 @@ lemma EndToEndTrace.loop_concat (etr₁ : EndToEndTrace (Graph.loop g) ρ₁ ρ
|
||||
simp only [Graph.loop_inputs, Graph.loop_outputs, List.mem_singleton] at h₁ h₂ k₁ k₂
|
||||
subst h₁; subst h₂; subst k₁; subst k₂
|
||||
exact ⟨g.loopIn, List.mem_singleton_self _, g.loopOut, List.mem_singleton_self _,
|
||||
Trace.concat tr₁ loop_edge_out_in tr₂⟩
|
||||
tr₁ ++< loop_edge_out_in >++ tr₂⟩
|
||||
|
||||
lemma EndToEndTrace.loop_empty {ρ : Env} : EndToEndTrace (Graph.loop g) ρ ρ := by
|
||||
/-- A loop can be executed zero times. -/
|
||||
noncomputable def EndToEndTrace.loop_empty {ρ : Env} : EndToEndTrace (Graph.loop g) ρ ρ := by
|
||||
have hedge : ((g.loopIn, g.loopOut) : (Graph.loop g).Edge) ∈ (Graph.loop g).edges :=
|
||||
List.mem_append_right _ (List.mem_cons_of_mem _ (List.mem_cons_self _ _))
|
||||
exact ⟨g.loopIn, List.mem_singleton_self _, g.loopOut, List.mem_singleton_self _,
|
||||
Trace.concat (Trace.single (loop_nodes_at_in ▸ EvalBasicStmts.nil)) hedge
|
||||
(Trace.single (loop_nodes_at_out ▸ EvalBasicStmts.nil))⟩
|
||||
Trace.single (loop_nodes_at_in ▸ EvalBasicStmtOpt.none) ++< hedge >++
|
||||
Trace.single (loop_nodes_at_out ▸ EvalBasicStmtOpt.none)⟩
|
||||
|
||||
end Loop
|
||||
|
||||
/-! ### Singletons, wrap, and the main result -/
|
||||
|
||||
lemma EndToEndTrace.singleton {bss : List BasicStmt} {ρ₁ ρ₂ : Env}
|
||||
(h : EvalBasicStmts ρ₁ bss ρ₂) : EndToEndTrace (Graph.singleton bss) ρ₁ ρ₂ :=
|
||||
/-- A CFG consisting of only a single node has a trace through it corresponding to that node. -/
|
||||
noncomputable def EndToEndTrace.singleton {o : Option BasicStmt} {ρ₁ ρ₂ : Env}
|
||||
(h : EvalBasicStmtOpt ρ₁ o ρ₂) : EndToEndTrace (Graph.singleton o) ρ₁ ρ₂ :=
|
||||
⟨(0 : Fin 1), List.mem_singleton_self _, (0 : Fin 1), List.mem_singleton_self _,
|
||||
Trace.single h⟩
|
||||
|
||||
lemma EndToEndTrace.singleton_nil (ρ : Env) :
|
||||
EndToEndTrace (Graph.singleton []) ρ ρ :=
|
||||
EndToEndTrace.singleton EvalBasicStmts.nil
|
||||
/-- If a CFG's only node is empty, the no-op trace exists through it. -/
|
||||
noncomputable def EndToEndTrace.singleton_nil (ρ : Env) :
|
||||
EndToEndTrace (Graph.singleton none) ρ ρ :=
|
||||
EndToEndTrace.singleton EvalBasicStmtOpt.none
|
||||
|
||||
lemma EndToEndTrace.wrap {g : Graph} {ρ₁ ρ₂ : Env}
|
||||
/-- Invoking 'Graph.wrap` (which ensures a single entry and exit node for a CFG)
|
||||
does not invalidate traces in the original graph. -/
|
||||
noncomputable def EndToEndTrace.wrap {g : Graph} {ρ₁ ρ₂ : Env}
|
||||
(etr : EndToEndTrace g ρ₁ ρ₂) : EndToEndTrace (Graph.wrap g) ρ₁ ρ₂ :=
|
||||
(EndToEndTrace.singleton_nil ρ₁).concat (etr.concat (EndToEndTrace.singleton_nil ρ₂))
|
||||
|
||||
theorem Stmt.cfg_sufficient {s : Stmt} {ρ₁ ρ₂ : Env}
|
||||
/-- Key result: the control flow graph admits every execution that's made
|
||||
possible by a language's semantics. Thus, the CFG encodes _at least_ all
|
||||
semantically-possible executions. Informally, we can conclude from this
|
||||
that if we compute a result that using the graph's edges to determine
|
||||
what's possible, this result will not disagree with the semantics.
|
||||
|
||||
Note that a CFG like $K_4$ (where the nodes are basic blocks) is
|
||||
technically also a sufficient graph, but is very likely meaningless in that
|
||||
it grossly overestimates the possible execution paths in the language, and
|
||||
thus is bound to produce less-than-specific results. There is as yet no
|
||||
result in this framework that the CFG we produce is _minimal_: loosely,
|
||||
posessing only edges for things that are admitted by the semantics.
|
||||
This is difficult to state (in its strongest form, this would
|
||||
require the CFG to be able to detect something like `while (alwaysFalse)`,
|
||||
and so remains a TODO. -/
|
||||
noncomputable def Stmt.cfg_sufficient {s : Stmt} {ρ₁ ρ₂ : Env}
|
||||
(h : EvalStmt ρ₁ s ρ₂) : EndToEndTrace s.cfg ρ₁ ρ₂ := by
|
||||
induction h with
|
||||
| basic ρ₁ ρ₂ bs hbs =>
|
||||
exact EndToEndTrace.singleton (EvalBasicStmts.cons hbs EvalBasicStmts.nil)
|
||||
exact EndToEndTrace.singleton (EvalBasicStmtOpt.some hbs)
|
||||
| andThen ρ₁ ρ₂ ρ₃ s₁ s₂ _ _ ih₁ ih₂ =>
|
||||
exact ih₁.concat ih₂
|
||||
| ifTrue ρ₁ ρ₂ e z s₁ s₂ _ _ _ ih =>
|
||||
@@ -190,24 +246,28 @@ theorem Stmt.cfg_sufficient {s : Stmt} {ρ₁ ρ₂ : Env}
|
||||
| whileFalse ρ e s _ =>
|
||||
exact EndToEndTrace.loop_empty
|
||||
|
||||
/-! ### The wrapped graph's entry has no predecessors (Agda's "ugly" block) -/
|
||||
|
||||
/-- The input / entry node generated by `Graph.wrap`. -/
|
||||
def Graph.wrapInput (g : Graph) : (Graph.wrap g).Index :=
|
||||
(0 : Fin 1).castAdd ((g ⤳ Graph.singleton []).size)
|
||||
(0 : Fin 1).castAdd ((g ⤳ Graph.singleton none).size)
|
||||
|
||||
/-- The output / exit node generated by `Graph.wrap`. -/
|
||||
def Graph.wrapOutput (g : Graph) : (Graph.wrap g).Index :=
|
||||
Fin.natAdd 1 ((Fin.natAdd g.size (0 : Fin 1)))
|
||||
|
||||
/-- The `Graph.wrapInput` is, indeed, the graph's only input after `Graph.wrap`. -/
|
||||
lemma Graph.wrap_inputs (g : Graph) :
|
||||
(Graph.wrap g).inputs = [g.wrapInput] := rfl
|
||||
|
||||
/-- The `Graph.wrapInput` is, indeed, the graph's only output after `Graph.wrap`. -/
|
||||
lemma Graph.wrap_outputs (g : Graph) :
|
||||
(Graph.wrap g).outputs = [g.wrapOutput] := rfl
|
||||
|
||||
/-- When sequencing (proven here with `Graph.singleton` on the left), no edges
|
||||
exist from the right-hand graph back to the left. -/
|
||||
private lemma not_mem_edges_castAdd_sequence {g₂ : Graph} (i : Fin 1)
|
||||
(idx : (Graph.singleton [] ⤳ g₂).Index) :
|
||||
((idx, i.castAdd g₂.size) : (Graph.singleton [] ⤳ g₂).Edge)
|
||||
∉ (Graph.singleton [] ⤳ g₂).edges := by
|
||||
(idx : (Graph.singleton none ⤳ g₂).Index) :
|
||||
((idx, i.castAdd g₂.size) : (Graph.singleton none ⤳ g₂).Edge)
|
||||
∉ (Graph.singleton none ⤳ g₂).edges := by
|
||||
intro h
|
||||
rcases List.mem_append.mp h with h' | h'
|
||||
· rcases List.mem_append.mp h' with h'' | h''
|
||||
@@ -221,6 +281,7 @@ private lemma not_mem_edges_castAdd_sequence {g₂ : Graph} (i : Fin 1)
|
||||
obtain ⟨j, -, heq⟩ := List.mem_map.mp hb
|
||||
exact Fin.castAdd_ne_natAdd i j heq.symm
|
||||
|
||||
/-- The input node of a graph after `Graph.wrap` has no predecessors. -/
|
||||
lemma Graph.wrap_predecessors_eq_nil (g : Graph) (idx : (Graph.wrap g).Index)
|
||||
(h : idx ∈ (Graph.wrap g).inputs) :
|
||||
(Graph.wrap g).predecessors idx = [] := by
|
||||
@@ -228,6 +289,6 @@ lemma Graph.wrap_predecessors_eq_nil (g : Graph) (idx : (Graph.wrap g).Index)
|
||||
subst h
|
||||
rw [GGraph.predecessors, List.filter_eq_nil_iff]
|
||||
intro idx' _
|
||||
simpa using not_mem_edges_castAdd_sequence (g₂ := g ⤳ Graph.singleton []) 0 idx'
|
||||
simpa using not_mem_edges_castAdd_sequence (g₂ := g ⤳ Graph.singleton none) 0 idx'
|
||||
|
||||
end Spa
|
||||
|
||||
@@ -46,22 +46,22 @@ inductive EvalExpr : Env → Expr → Value → Prop
|
||||
/-- Inference rules for evaluating a basic statement (`Spa.BasicStmt`) in
|
||||
a given environment, potentially changing the environment.
|
||||
Pretty standard big-step evaluation. -/
|
||||
inductive EvalBasicStmt : Env → BasicStmt → Env → Prop
|
||||
inductive EvalBasicStmt : Env → BasicStmt → Env → Type
|
||||
| noop (ρ : Env) : EvalBasicStmt ρ .noop ρ
|
||||
| assign (ρ : Env) (x : String) (e : Expr) (v : Value) :
|
||||
EvalExpr ρ e v → EvalBasicStmt ρ (.assign x e) ((x, v) :: ρ)
|
||||
|
||||
/-- Inference rules for evaluating a sequence of basic statements. -/
|
||||
inductive EvalBasicStmts : Env → List BasicStmt → Env → Prop
|
||||
| nil {ρ : Env} : EvalBasicStmts ρ [] ρ
|
||||
| cons {ρ₁ ρ₂ ρ₃ : Env} {bs : BasicStmt} {bss : List BasicStmt} :
|
||||
EvalBasicStmt ρ₁ bs ρ₂ → EvalBasicStmts ρ₂ bss ρ₃ →
|
||||
EvalBasicStmts ρ₁ (bs :: bss) ρ₃
|
||||
/-- Inference rules for evaluating a basic-statement-or-nothing,
|
||||
which is the current representation of CFGs nodes. -/
|
||||
inductive EvalBasicStmtOpt : Env → Option BasicStmt → Env → Type
|
||||
| none {ρ : Env} : EvalBasicStmtOpt ρ Option.none ρ
|
||||
| some {ρ₁ ρ₂ : Env} {bs : BasicStmt} :
|
||||
EvalBasicStmt ρ₁ bs ρ₂ → EvalBasicStmtOpt ρ₁ (Option.some bs) ρ₂
|
||||
|
||||
/-- Inference rules for evaluating statements (`Spa.Stmt`) in a given
|
||||
environment, potentially changing the environment.
|
||||
Pretty standard big-step evaluation. -/
|
||||
inductive EvalStmt : Env → Stmt → Env → Prop
|
||||
inductive EvalStmt : Env → Stmt → Env → Type
|
||||
| basic (ρ₁ ρ₂ : Env) (bs : BasicStmt) :
|
||||
EvalBasicStmt ρ₁ bs ρ₂ → EvalStmt ρ₁ (.basic bs) ρ₂
|
||||
| andThen (ρ₁ ρ₂ ρ₃ : Env) (s₁ s₂ : Stmt) :
|
||||
|
||||
18
lean/Spa/Language/Tagged/Basic.lean
Normal file
18
lean/Spa/Language/Tagged/Basic.lean
Normal file
@@ -0,0 +1,18 @@
|
||||
import Spa.Language.Base
|
||||
import Spa.Language.Tagged.Id
|
||||
import Spa.Language.Tagged.Derive
|
||||
|
||||
derive_tagged Spa.Expr Spa.BasicStmt Spa.Stmt
|
||||
|
||||
namespace Spa
|
||||
|
||||
def tagStmt (s : Stmt) : Stmt.Tagged RawId := (s.tag 0).1
|
||||
|
||||
def Stmt.Tagged.subtreeIds {τ : Type} (s : Stmt.Tagged τ) : List τ :=
|
||||
s.foldTags (· :: ·) []
|
||||
|
||||
def Stmt.Tagged.isInLoopBody {τ : Type} [DecidableEq τ]
|
||||
(body : Stmt.Tagged τ) (id : τ) : Bool :=
|
||||
decide (id ∈ body.subtreeIds)
|
||||
|
||||
end Spa
|
||||
417
lean/Spa/Language/Tagged/DESCENDANT-TRACKING.md
Normal file
417
lean/Spa/Language/Tagged/DESCENDANT-TRACKING.md
Normal file
@@ -0,0 +1,417 @@
|
||||
# Descendant tracking (parked)
|
||||
|
||||
This is the formally-verified **interval-labeling / descendant** machinery that
|
||||
used to live in `Id.lean` and `Properties.lean`. It let you decide "is node `a`
|
||||
a descendant of node `b`?" with two integer comparisons on their identifiers,
|
||||
and *proved* that numeric test equivalent to structural subtree containment.
|
||||
|
||||
It was removed because the descendant test is a *computational optimization*:
|
||||
the same question can be answered by walking the AST, and nothing in the current
|
||||
pipeline needs the fast test yet. The proofs (a rose-tree flattening + a
|
||||
postorder `Good` invariant) are a real mechanization cost to carry. Parked here
|
||||
so it can be restored verbatim when LICM actually wants it.
|
||||
|
||||
## What stays in the live code
|
||||
|
||||
- `NodeId` collapses to a single unique index (`{ post : ℕ }`); `tag` still
|
||||
assigns each node a distinct postorder number.
|
||||
- The bidirectional mapping (`erase`/`tag` + `erase_tagStmt`) stays in
|
||||
`Properties.lean`.
|
||||
- The labelled-CFG id↔state mapping (`Cfg.lean`) is independent of this and is
|
||||
unaffected.
|
||||
|
||||
## Revival checklist
|
||||
|
||||
1. In `Id.lean`, give `NodeId` back its descendant-count field and the test:
|
||||
|
||||
```lean
|
||||
structure NodeId where
|
||||
post : ℕ
|
||||
desc : ℕ -- number of proper descendants (subtree size − 1); leaf = 0
|
||||
deriving DecidableEq, Repr
|
||||
|
||||
namespace NodeId
|
||||
|
||||
/-- Left endpoint of the node's postorder interval `[lo, post]`. -/
|
||||
def lo (a : NodeId) : ℕ := a.post - a.desc
|
||||
|
||||
/-- `a` is a descendant-or-self of `b`: `a.post` lies in `b`'s interval. -/
|
||||
def DescendantOf (a b : NodeId) : Prop := b.lo ≤ a.post ∧ a.post ≤ b.post
|
||||
|
||||
instance (a b : NodeId) : Decidable (DescendantOf a b) := by
|
||||
unfold DescendantOf; infer_instance
|
||||
|
||||
end NodeId
|
||||
```
|
||||
|
||||
2. In `Derive.lean`, make the generated `tag` store the descendant count again:
|
||||
change the emitted identifier in `mkTag` from `(⟨$last⟩ : $nId)` back to
|
||||
`(⟨$last, $last - n⟩ : $nId)`.
|
||||
|
||||
3. Paste the Lean block below back into `Properties.lean` (after the round-trip
|
||||
theorems). It builds against the `id.lo = lo`-premise form of `Good` and the
|
||||
childcount (`desc`) identifier. The headline result is
|
||||
`descendant_iff_tagStmt`; everything else is supporting machinery.
|
||||
|
||||
## The parked proofs
|
||||
|
||||
```lean
|
||||
/-- A rose tree of identifiers: the uniform shape underlying all three tagged
|
||||
AST types, used to reason about the postorder labeling generically. -/
|
||||
inductive IdTree where
|
||||
| node (id : NodeId) (children : List IdTree)
|
||||
|
||||
namespace IdTree
|
||||
|
||||
def rootId : IdTree → NodeId
|
||||
| .node id _ => id
|
||||
|
||||
@[simp] theorem rootId_node (id : NodeId) (cs : List IdTree) :
|
||||
(IdTree.node id cs).rootId = id := rfl
|
||||
|
||||
mutual
|
||||
def subtrees : IdTree → List IdTree
|
||||
| .node id cs => .node id cs :: subtreesList cs
|
||||
def subtreesList : List IdTree → List IdTree
|
||||
| [] => []
|
||||
| c :: cs => subtrees c ++ subtreesList cs
|
||||
end
|
||||
|
||||
@[simp] theorem subtrees_node (id : NodeId) (cs : List IdTree) :
|
||||
subtrees (.node id cs) = .node id cs :: subtreesList cs := rfl
|
||||
|
||||
@[simp] theorem subtreesList_nil : subtreesList [] = [] := rfl
|
||||
|
||||
@[simp] theorem subtreesList_cons (c : IdTree) (cs : List IdTree) :
|
||||
subtreesList (c :: cs) = subtrees c ++ subtreesList cs := rfl
|
||||
|
||||
def posts (t : IdTree) : List ℕ := (subtrees t).map (fun s => s.rootId.post)
|
||||
|
||||
def postsList (cs : List IdTree) : List ℕ := (subtreesList cs).map (fun s => s.rootId.post)
|
||||
|
||||
@[simp] theorem posts_node (id : NodeId) (cs : List IdTree) :
|
||||
posts (.node id cs) = id.post :: postsList cs := rfl
|
||||
|
||||
@[simp] theorem postsList_nil : postsList [] = [] := rfl
|
||||
|
||||
@[simp] theorem postsList_cons (c : IdTree) (cs : List IdTree) :
|
||||
postsList (c :: cs) = posts c ++ postsList cs := by
|
||||
simp [posts, postsList]
|
||||
|
||||
end IdTree
|
||||
|
||||
def Expr.Tagged.toIdTree : Expr.Tagged NodeId → IdTree
|
||||
| .add t a b => .node t [a.toIdTree, b.toIdTree]
|
||||
| .sub t a b => .node t [a.toIdTree, b.toIdTree]
|
||||
| .var t _ => .node t []
|
||||
| .num t _ => .node t []
|
||||
|
||||
def BasicStmt.Tagged.toIdTree : BasicStmt.Tagged NodeId → IdTree
|
||||
| .assign t _ e => .node t [e.toIdTree]
|
||||
| .noop t => .node t []
|
||||
|
||||
def Stmt.Tagged.toIdTree : Stmt.Tagged NodeId → IdTree
|
||||
| .basic t bs => .node t [bs.toIdTree]
|
||||
| .andThen t a b => .node t [a.toIdTree, b.toIdTree]
|
||||
| .ifElse t e a b => .node t [e.toIdTree, a.toIdTree, b.toIdTree]
|
||||
| .whileLoop t e s => .node t [e.toIdTree, s.toIdTree]
|
||||
|
||||
mutual
|
||||
inductive Good : ℕ → IdTree → Prop
|
||||
| mk {lo : ℕ} {id : NodeId} {cs : List IdTree} :
|
||||
id.lo = lo → GoodChildren lo cs id.post →
|
||||
Good lo (.node id cs)
|
||||
inductive GoodChildren : ℕ → List IdTree → ℕ → Prop
|
||||
| nil {pos : ℕ} : GoodChildren pos [] pos
|
||||
| cons {cur : ℕ} {c : IdTree} {cs : List IdTree} {pos : ℕ} :
|
||||
Good cur c → GoodChildren (c.rootId.post + 1) cs pos →
|
||||
GoodChildren cur (c :: cs) pos
|
||||
end
|
||||
|
||||
theorem Good.lo_le_post {lo : ℕ} {t : IdTree} (h : Good lo t) : lo ≤ t.rootId.post := by
|
||||
cases h with
|
||||
| mk hlo _ => simp only [NodeId.lo] at hlo; simp only [IdTree.rootId_node]; omega
|
||||
|
||||
theorem GoodChildren.cur_le_pos : ∀ {cur : ℕ} (cs : List IdTree) {pos : ℕ},
|
||||
GoodChildren cur cs pos → cur ≤ pos
|
||||
| _, [], _, h => by cases h; exact le_rfl
|
||||
| _, c :: cs, _, h => by
|
||||
cases h with
|
||||
| cons hc hcs =>
|
||||
have := hc.lo_le_post
|
||||
have := GoodChildren.cur_le_pos cs hcs
|
||||
omega
|
||||
|
||||
mutual
|
||||
theorem Good.mem_posts : ∀ {lo : ℕ} (t : IdTree), Good lo t →
|
||||
∀ x, x ∈ IdTree.posts t ↔ lo ≤ x ∧ x ≤ t.rootId.post
|
||||
| _, .node id cs, h, x => by
|
||||
cases h with
|
||||
| mk hlo hch =>
|
||||
simp only [IdTree.posts_node, List.mem_cons, IdTree.rootId_node]
|
||||
rw [GoodChildren.mem_postsList cs hch x]
|
||||
simp only [NodeId.lo] at hlo
|
||||
omega
|
||||
theorem GoodChildren.mem_postsList : ∀ {cur : ℕ} (cs : List IdTree) {pos : ℕ},
|
||||
GoodChildren cur cs pos → ∀ x, x ∈ IdTree.postsList cs ↔ cur ≤ x ∧ x < pos
|
||||
| _, [], _, h, x => by
|
||||
cases h
|
||||
simp only [IdTree.postsList_nil]
|
||||
constructor
|
||||
· intro hx; exact absurd hx (List.not_mem_nil x)
|
||||
· rintro ⟨h1, h2⟩; exfalso; omega
|
||||
| _, c :: cs, _, h, x => by
|
||||
cases h with
|
||||
| cons hc hcs =>
|
||||
simp only [IdTree.postsList_cons, List.mem_append]
|
||||
rw [Good.mem_posts c hc x, GoodChildren.mem_postsList cs hcs x]
|
||||
have := hc.lo_le_post
|
||||
have := GoodChildren.cur_le_pos cs hcs
|
||||
omega
|
||||
end
|
||||
|
||||
mutual
|
||||
theorem Good.nodup_posts : ∀ {lo : ℕ} (t : IdTree), Good lo t → (IdTree.posts t).Nodup
|
||||
| _, .node id cs, h => by
|
||||
cases h with
|
||||
| mk hlo hch =>
|
||||
simp only [IdTree.posts_node, List.nodup_cons]
|
||||
refine ⟨?_, GoodChildren.nodup_postsList cs hch⟩
|
||||
intro hmem
|
||||
rw [GoodChildren.mem_postsList cs hch id.post] at hmem
|
||||
omega
|
||||
theorem GoodChildren.nodup_postsList : ∀ {cur : ℕ} (cs : List IdTree) {pos : ℕ},
|
||||
GoodChildren cur cs pos → (IdTree.postsList cs).Nodup
|
||||
| _, [], _, h => by cases h; simp only [IdTree.postsList_nil, List.nodup_nil]
|
||||
| _, c :: cs, _, h => by
|
||||
cases h with
|
||||
| cons hc hcs =>
|
||||
simp only [IdTree.postsList_cons, List.nodup_append]
|
||||
refine ⟨Good.nodup_posts c hc, GoodChildren.nodup_postsList cs hcs, ?_⟩
|
||||
intro x hx1 hx2
|
||||
rw [Good.mem_posts c hc x] at hx1
|
||||
rw [GoodChildren.mem_postsList cs hcs x] at hx2
|
||||
omega
|
||||
end
|
||||
|
||||
mutual
|
||||
theorem Good.subtree_good : ∀ {lo : ℕ} (t : IdTree), Good lo t →
|
||||
∀ s ∈ IdTree.subtrees t, Good s.rootId.lo s
|
||||
| _, .node id cs, h, s, hs => by
|
||||
cases h with
|
||||
| mk hlo hch =>
|
||||
rw [IdTree.subtrees_node, List.mem_cons] at hs
|
||||
rcases hs with rfl | hs
|
||||
· simp only [IdTree.rootId_node]; rw [hlo]; exact Good.mk hlo hch
|
||||
· exact GoodChildren.subtree_good cs hch s hs
|
||||
theorem GoodChildren.subtree_good : ∀ {cur : ℕ} (cs : List IdTree) {pos : ℕ},
|
||||
GoodChildren cur cs pos → ∀ s ∈ IdTree.subtreesList cs, Good s.rootId.lo s
|
||||
| _, [], _, _, s, hs => by simp only [IdTree.subtreesList_nil, List.not_mem_nil] at hs
|
||||
| _, c :: cs, _, h, s, hs => by
|
||||
cases h with
|
||||
| cons hc hcs =>
|
||||
rw [IdTree.subtreesList_cons, List.mem_append] at hs
|
||||
rcases hs with hs | hs
|
||||
· exact Good.subtree_good c hc s hs
|
||||
· exact GoodChildren.subtree_good cs hcs s hs
|
||||
end
|
||||
|
||||
mutual
|
||||
theorem IdTree.subtrees_subset : ∀ (t : IdTree) {b : IdTree},
|
||||
b ∈ subtrees t → subtrees b ⊆ subtrees t
|
||||
| .node id cs, b, hb => by
|
||||
rw [subtrees_node, List.mem_cons] at hb
|
||||
rcases hb with rfl | hb
|
||||
· exact fun _ h => h
|
||||
· intro x hx
|
||||
rw [subtrees_node, List.mem_cons]
|
||||
exact Or.inr (IdTree.subtreesList_subset cs hb hx)
|
||||
theorem IdTree.subtreesList_subset : ∀ (cs : List IdTree) {b : IdTree},
|
||||
b ∈ subtreesList cs → subtrees b ⊆ subtreesList cs
|
||||
| [], b, hb => by simp only [subtreesList_nil, List.not_mem_nil] at hb
|
||||
| c :: cs, b, hb => by
|
||||
rw [subtreesList_cons, List.mem_append] at hb
|
||||
intro x hx
|
||||
rw [subtreesList_cons, List.mem_append]
|
||||
rcases hb with hb | hb
|
||||
· exact Or.inl (IdTree.subtrees_subset c hb hx)
|
||||
· exact Or.inr (IdTree.subtreesList_subset cs hb hx)
|
||||
end
|
||||
|
||||
theorem IdTree.eq_of_post_eq {l : List IdTree}
|
||||
(h : (l.map (fun s => s.rootId.post)).Nodup) {a c : IdTree}
|
||||
(ha : a ∈ l) (hc : c ∈ l) (hpost : a.rootId.post = c.rootId.post) : a = c := by
|
||||
induction l with
|
||||
| nil => exact absurd ha (List.not_mem_nil a)
|
||||
| cons d ds ih =>
|
||||
simp only [List.map_cons, List.nodup_cons] at h
|
||||
obtain ⟨hd, htl⟩ := h
|
||||
simp only [List.mem_cons] at ha hc
|
||||
rcases ha with rfl | ha <;> rcases hc with rfl | hc
|
||||
· rfl
|
||||
· exfalso; apply hd; rw [hpost]; exact List.mem_map_of_mem _ hc
|
||||
· exfalso; apply hd; rw [← hpost]; exact List.mem_map_of_mem _ ha
|
||||
· exact ih htl ha hc
|
||||
|
||||
theorem descendant_iff_of_good {lo : ℕ} {t : IdTree} (hg : Good lo t)
|
||||
{a b : IdTree} (ha : a ∈ IdTree.subtrees t) (hb : b ∈ IdTree.subtrees t) :
|
||||
a.rootId.DescendantOf b.rootId ↔ a ∈ IdTree.subtrees b := by
|
||||
have hgb : Good b.rootId.lo b := Good.subtree_good t hg b hb
|
||||
constructor
|
||||
· rintro ⟨h1, h2⟩
|
||||
have hmem : a.rootId.post ∈ IdTree.posts b := by
|
||||
rw [Good.mem_posts b hgb a.rootId.post]; exact ⟨h1, h2⟩
|
||||
rw [IdTree.posts, List.mem_map] at hmem
|
||||
obtain ⟨c, hc_mem, hc_post⟩ := hmem
|
||||
have hc_t : c ∈ IdTree.subtrees t := IdTree.subtrees_subset t hb hc_mem
|
||||
have hac : a = c :=
|
||||
IdTree.eq_of_post_eq (hg.nodup_posts t) ha hc_t hc_post.symm
|
||||
rw [hac]; exact hc_mem
|
||||
· intro hsub
|
||||
have hmem : a.rootId.post ∈ IdTree.posts b := by
|
||||
rw [IdTree.posts, List.mem_map]; exact ⟨a, hsub, rfl⟩
|
||||
rw [Good.mem_posts b hgb a.rootId.post] at hmem
|
||||
exact hmem
|
||||
|
||||
/-! ### Tagging produces a good tree
|
||||
|
||||
We bridge from the `tag` traversal to the abstract `Good` invariant, by induction
|
||||
on the plain AST. Each lemma also records that the returned counter is one past
|
||||
the root's postorder index. -/
|
||||
|
||||
theorem Expr.tag_spec : ∀ (e : Expr) (n : ℕ),
|
||||
Good n (e.tag n).1.toIdTree ∧ (e.tag n).1.toIdTree.rootId.post + 1 = (e.tag n).2 := by
|
||||
intro e
|
||||
induction e with
|
||||
| num k =>
|
||||
intro n
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree]
|
||||
exact Good.mk (by simp only [NodeId.lo]; omega) GoodChildren.nil
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| var x =>
|
||||
intro n
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree]
|
||||
exact Good.mk (by simp only [NodeId.lo]; omega) GoodChildren.nil
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| add a b iha ihb =>
|
||||
intro n
|
||||
obtain ⟨gA, pA⟩ := iha n
|
||||
obtain ⟨gB, pB⟩ := ihb (a.tag n).2
|
||||
have lA := gA.lo_le_post
|
||||
have lB := gB.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gA ?_
|
||||
rw [pA]; refine GoodChildren.cons gB ?_; rw [pB]; exact GoodChildren.nil
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| sub a b iha ihb =>
|
||||
intro n
|
||||
obtain ⟨gA, pA⟩ := iha n
|
||||
obtain ⟨gB, pB⟩ := ihb (a.tag n).2
|
||||
have lA := gA.lo_le_post
|
||||
have lB := gB.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gA ?_
|
||||
rw [pA]; refine GoodChildren.cons gB ?_; rw [pB]; exact GoodChildren.nil
|
||||
· simp only [Expr.tag, Expr.Tagged.toIdTree, IdTree.rootId_node]
|
||||
|
||||
theorem BasicStmt.tag_spec : ∀ (bs : BasicStmt) (n : ℕ),
|
||||
Good n (bs.tag n).1.toIdTree ∧ (bs.tag n).1.toIdTree.rootId.post + 1 = (bs.tag n).2 := by
|
||||
intro bs
|
||||
cases bs with
|
||||
| noop =>
|
||||
intro n
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [BasicStmt.tag, BasicStmt.Tagged.toIdTree]
|
||||
exact Good.mk (by simp only [NodeId.lo]; omega) GoodChildren.nil
|
||||
· simp only [BasicStmt.tag, BasicStmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| assign x e =>
|
||||
intro n
|
||||
obtain ⟨gE, pE⟩ := Expr.tag_spec e n
|
||||
have lE := gE.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [BasicStmt.tag, BasicStmt.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gE ?_
|
||||
rw [pE]; exact GoodChildren.nil
|
||||
· simp only [BasicStmt.tag, BasicStmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
|
||||
theorem Stmt.tag_spec : ∀ (s : Stmt) (n : ℕ),
|
||||
Good n (s.tag n).1.toIdTree ∧ (s.tag n).1.toIdTree.rootId.post + 1 = (s.tag n).2 := by
|
||||
intro s
|
||||
induction s with
|
||||
| basic bs =>
|
||||
intro n
|
||||
obtain ⟨gBs, pBs⟩ := BasicStmt.tag_spec bs n
|
||||
have lBs := gBs.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gBs ?_
|
||||
rw [pBs]; exact GoodChildren.nil
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| andThen a b iha ihb =>
|
||||
intro n
|
||||
obtain ⟨gA, pA⟩ := iha n
|
||||
obtain ⟨gB, pB⟩ := ihb (a.tag n).2
|
||||
have lA := gA.lo_le_post
|
||||
have lB := gB.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gA ?_
|
||||
rw [pA]; refine GoodChildren.cons gB ?_; rw [pB]; exact GoodChildren.nil
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| ifElse e a b iha ihb =>
|
||||
intro n
|
||||
obtain ⟨gE, pE⟩ := Expr.tag_spec e n
|
||||
obtain ⟨gA, pA⟩ := iha (e.tag n).2
|
||||
obtain ⟨gB, pB⟩ := ihb (a.tag (e.tag n).2).2
|
||||
have lE := gE.lo_le_post
|
||||
have lA := gA.lo_le_post
|
||||
have lB := gB.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gE ?_
|
||||
rw [pE]; refine GoodChildren.cons gA ?_
|
||||
rw [pA]; refine GoodChildren.cons gB ?_; rw [pB]; exact GoodChildren.nil
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
| whileLoop e s ih =>
|
||||
intro n
|
||||
obtain ⟨gE, pE⟩ := Expr.tag_spec e n
|
||||
obtain ⟨gS, pS⟩ := ih (e.tag n).2
|
||||
have lE := gE.lo_le_post
|
||||
have lS := gS.lo_le_post
|
||||
refine ⟨?_, ?_⟩
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree]
|
||||
refine Good.mk ?_ ?_
|
||||
· simp only [NodeId.lo]; omega
|
||||
· refine GoodChildren.cons gE ?_
|
||||
rw [pE]; refine GoodChildren.cons gS ?_; rw [pS]; exact GoodChildren.nil
|
||||
· simp only [Stmt.tag, Stmt.Tagged.toIdTree, IdTree.rootId_node]
|
||||
|
||||
/-- A freshly tagged program is a well-tagged tree (rooted at postorder start `0`). -/
|
||||
theorem good_tagStmt (s : Stmt) : Good 0 (tagStmt s).toIdTree :=
|
||||
(Stmt.tag_spec s 0).1
|
||||
|
||||
/-- **Descendant characterization.** The numeric `NodeId.DescendantOf` relation on
|
||||
two nodes of a tagged program holds exactly when one is structurally contained in
|
||||
the other's subtree. -/
|
||||
theorem descendant_iff_tagStmt (s : Stmt) {a b : IdTree}
|
||||
(ha : a ∈ IdTree.subtrees (tagStmt s).toIdTree)
|
||||
(hb : b ∈ IdTree.subtrees (tagStmt s).toIdTree) :
|
||||
a.rootId.DescendantOf b.rootId ↔ a ∈ IdTree.subtrees b :=
|
||||
descendant_iff_of_good (good_tagStmt s) ha hb
|
||||
```
|
||||
509
lean/Spa/Language/Tagged/Derive.lean
Normal file
509
lean/Spa/Language/Tagged/Derive.lean
Normal file
@@ -0,0 +1,509 @@
|
||||
import Lean
|
||||
import Mathlib.Tactic.DeriveTraversable
|
||||
import Spa.Language.Base
|
||||
import Spa.Language.Tagged.Id
|
||||
|
||||
/-!
|
||||
# The `derive_tagged` command
|
||||
|
||||
`derive_tagged T₁ T₂ … Tₙ` takes a family of (possibly mutually recursive)
|
||||
inductive types and generates, for each `Tᵢ`:
|
||||
|
||||
* a *tagged* mirror inductive `Tᵢ.Tagged (τ : Type)`, in which every constructor
|
||||
carries a leading `tag : τ` field and every field whose type is a family
|
||||
member is retyped to its `.Tagged τ` counterpart;
|
||||
* `Tᵢ.Tagged.erase : Tᵢ.Tagged τ → Tᵢ`, forgetting all tags;
|
||||
* `Tᵢ.tag : Tᵢ → ℕ → Tᵢ.Tagged RawId × ℕ`, assigning every node a unique
|
||||
`RawId` (its postorder index) by a single unified traversal that threads a
|
||||
counter; the whole family shares one counter, so identifiers are unique across
|
||||
types.
|
||||
|
||||
The generated declarations have exactly the shape of the hand-written reference;
|
||||
see `Spa/Language/Tagged/Basic.lean` (which invokes this command) and the proofs
|
||||
in `Spa/Language/Tagged/Properties.lean`.
|
||||
|
||||
Scope: the generator handles non-indexed inductives whose constructor fields are
|
||||
either scalars or *direct* references to a family member (which covers the object
|
||||
language). Nested occurrences such as `List Tᵢ` are not supported.
|
||||
-/
|
||||
|
||||
open Lean Elab Command Meta
|
||||
|
||||
namespace Spa.DeriveTagged
|
||||
|
||||
/-- One constructor field, classified as a recursive family reference or a scalar
|
||||
(whose type syntax we keep verbatim for the mirror inductive). -/
|
||||
structure FieldData where
|
||||
isRec : Bool
|
||||
recType : Name
|
||||
typeStx : Term
|
||||
|
||||
/-- A constructor: its original (full) name, short name, and fields. -/
|
||||
structure CtorData where
|
||||
origName : Name
|
||||
shortName : Name
|
||||
fields : Array FieldData
|
||||
|
||||
/-- A family member together with its constructors. -/
|
||||
structure TypeData where
|
||||
name : Name
|
||||
ctors : Array CtorData
|
||||
|
||||
def taggedOf (n : Name) : Name := n ++ `Tagged
|
||||
def eraseOf (n : Name) : Name := n ++ `Tagged ++ `erase
|
||||
def rootTagOf (n : Name) : Name := n ++ `Tagged ++ `rootTag
|
||||
def tagOf (n : Name) : Name := n ++ `tag
|
||||
def foldTagsOf (n : Name) : Name := n ++ `Tagged ++ `foldTags
|
||||
def wfOf (n : Name) : Name := n ++ `Tagged ++ `WF
|
||||
def narrowOf (n : Name) : Name := n ++ `Tagged ++ `narrow
|
||||
def narrowEraseOf (n : Name) : Name := n ++ `Tagged ++ `narrow_erase
|
||||
def tagLeOf (n : Name) : Name := n ++ `tag_le
|
||||
def tagRootTagPostOf (n : Name) : Name := n ++ `tag_rootTag_post
|
||||
def tagWfOf (n : Name) : Name := n ++ `tag_wf
|
||||
|
||||
/-- Project the `i`-th conjunct (1-based) out of `hyp`, which has type a
|
||||
right-nested `And` of `total` conjuncts, e.g. `hyp |>.2 |>.2 |>.1`. -/
|
||||
def projAnd {m : Type → Type} [Monad m] [MonadQuotation m]
|
||||
(hyp : Term) (i total : Nat) : m Term := do
|
||||
let mut t := hyp
|
||||
for _ in [0:i-1] do
|
||||
t ← `($t |>.2)
|
||||
if i < total then
|
||||
t ← `($t |>.1)
|
||||
return t
|
||||
|
||||
/-- Combine a non-empty array of propositions into a right-nested conjunction. -/
|
||||
def mkAndR {m : Type → Type} [Monad m] [MonadQuotation m]
|
||||
(cs : Array Term) : m Term := do
|
||||
let mut t := cs.back!
|
||||
for c in cs.pop.reverse do
|
||||
t ← `($c ∧ $t)
|
||||
return t
|
||||
|
||||
/-- For a constructor, return one entry per *recursive* field: its argument
|
||||
identifier, the family member it references, and the start-counter expression at
|
||||
which it is tagged (`n`, then `(a.tag n).2`, …) — the same threading `mkTag`
|
||||
uses. -/
|
||||
def recChildren (cd : CtorData) (argNames : Array Ident) (nStart : Term) :
|
||||
CommandElabM (Array (Ident × Name × Term)) := do
|
||||
let mut res : Array (Ident × Name × Term) := #[]
|
||||
let mut cur := nStart
|
||||
for (f, a) in cd.fields.zip argNames do
|
||||
if f.isRec then
|
||||
res := res.push (a, f.recType, cur)
|
||||
cur ← `(($(mkIdent (tagOf f.recType)) $a $cur) |>.2)
|
||||
return res
|
||||
|
||||
/-- Inspect the family, classifying each constructor field. -/
|
||||
def gather (family : Array Name) (τ : Ident) : TermElabM (Array TypeData) := do
|
||||
let famSet : NameSet := family.foldl (·.insert ·) {}
|
||||
family.mapM fun tn => do
|
||||
let iv ← getConstInfoInduct tn
|
||||
let ctors ← iv.ctors.toArray.mapM fun cn => do
|
||||
let cv ← getConstInfoCtor cn
|
||||
let fields ← forallTelescopeReducing cv.type fun args _ => do
|
||||
let fieldArgs := args.extract iv.numParams args.size
|
||||
fieldArgs.mapM fun a => do
|
||||
let ty ← inferType a
|
||||
match ty.getAppFn.constName? with
|
||||
| some hn =>
|
||||
if famSet.contains hn then
|
||||
return { isRec := true, recType := hn, typeStx := ← `($(mkIdent (taggedOf hn)) $τ) }
|
||||
else
|
||||
return { isRec := false, recType := default, typeStx := ← Lean.PrettyPrinter.delab ty }
|
||||
| none =>
|
||||
return { isRec := false, recType := default, typeStx := ← Lean.PrettyPrinter.delab ty }
|
||||
return { origName := cn, shortName := cn.componentsRev.head!, fields }
|
||||
return { name := tn, ctors }
|
||||
|
||||
/-- The arrow type `τ → <fields…> → Self τ` of a tagged constructor. -/
|
||||
def ctorArrow (cd : CtorData) (self : Term) (τ : Ident) : TermElabM Term := do
|
||||
let mut t := self
|
||||
for f in cd.fields.reverse do
|
||||
t ← `($(f.typeStx) → $t)
|
||||
`($τ → $t)
|
||||
|
||||
/-- The tagged mirror inductives, one per family member. The family is a DAG
|
||||
(`Expr ← BasicStmt ← Stmt`), not genuinely mutual, so they are emitted as
|
||||
separate inductives in dependency order rather than a `mutual` block.
|
||||
|
||||
`Functor`/`Traversable` instances are derived separately by `mkDeriveInstances`
|
||||
below rather than via an inline `deriving` clause. -/
|
||||
def mkInductives (tds : Array TypeData) (τ : Ident) :
|
||||
CommandElabM (Array (TSyntax `command)) := do
|
||||
tds.mapM fun td => do
|
||||
let self ← `($(mkIdent (taggedOf td.name)) $τ)
|
||||
let ctors ← td.ctors.mapM fun cd => do
|
||||
let aty ← Command.liftTermElabM (ctorArrow cd self τ)
|
||||
`(Lean.Parser.Command.ctor| | $(mkIdent cd.shortName):ident : $aty)
|
||||
`(command| inductive $(mkIdent (taggedOf td.name)):ident ($τ : Type) where $ctors*)
|
||||
|
||||
/-- A `deriving instance Functor, Traversable for Tᵢ.Tagged` command per family
|
||||
member. Since every tagged type is a single-parameter, direct-recursive
|
||||
inductive in `τ`, Mathlib's deriving handler produces clean (`sorry`-free)
|
||||
instances, giving `map`, `traverse`, and the `Traversable.foldr`/`toList` folds
|
||||
for free.
|
||||
|
||||
These are emitted as *separate* commands in dependency order (rather than an
|
||||
inline `deriving` clause on each inductive) for two reasons: deriving
|
||||
`Stmt.Tagged` needs the `Expr.Tagged`/`BasicStmt.Tagged` instances already in
|
||||
scope, and — because every member's type name ends in `.Tagged` — the handler's
|
||||
auto-generated instance name (`instFunctorTagged`, built from the type's last
|
||||
component) collides across the family unless each derive sees the environment
|
||||
the previous one updated; separate commands give it that, so the names
|
||||
disambiguate to `instFunctorTagged`, `instFunctorTagged_1`, ….
|
||||
|
||||
The hand-written `foldTags` is retained alongside these: it is a
|
||||
structural-recursion fold that `simp`/`decide` reduce cleanly, unlike the
|
||||
abstract `Traversable.foldr` (defined via the `FreeMonoid`/`Const` applicative),
|
||||
which reduces under `decide`/`rfl` but not naive `simp` unfolding. -/
|
||||
def mkDeriveInstances (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
tds.mapM fun td =>
|
||||
`(command| deriving instance Functor, Traversable for $(mkIdent (taggedOf td.name)))
|
||||
|
||||
/-- The `erase` functions, one per family member (separate defs in dependency
|
||||
order — each calls only already-defined lower members). -/
|
||||
def mkErase (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
tds.mapM fun td => do
|
||||
let mut pats : Array Term := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map (fun i => mkIdent (.mkSimple s!"a{i}"))
|
||||
let pat ← `($(mkIdent (taggedOf td.name ++ cd.shortName)) _ $argNames*)
|
||||
let eraseArgs ← (cd.fields.zip argNames).mapM fun (f, a) =>
|
||||
if f.isRec then `($(mkIdent (eraseOf f.recType)) $a) else pure a
|
||||
let rhs ← `($(mkIdent cd.origName) $eraseArgs*)
|
||||
pats := pats.push pat
|
||||
rhss := rhss.push rhs
|
||||
`(command| def $(mkIdent (eraseOf td.name)) {τ : Type} :
|
||||
$(mkIdent (taggedOf td.name)) τ → $(mkIdent td.name) :=
|
||||
fun x => match x with $[| $pats => $rhss]*)
|
||||
|
||||
/-- The `rootTag` accessors (one non-recursive `def` per type). -/
|
||||
def mkRootTag (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let tIdent := mkIdent `t
|
||||
tds.mapM fun td => do
|
||||
let mut pats : Array Term := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let hole ← `(_)
|
||||
let wilds := Array.mkArray cd.fields.size hole
|
||||
pats := pats.push (← `($(mkIdent (taggedOf td.name ++ cd.shortName)) $tIdent $wilds*))
|
||||
rhss := rhss.push tIdent
|
||||
`(command| def $(mkIdent (rootTagOf td.name)) {τ : Type} :
|
||||
$(mkIdent (taggedOf td.name)) τ → τ :=
|
||||
fun x => match x with $[| $pats => $rhss]*)
|
||||
|
||||
/-- The postorder `tag` functions, one per family member (separate defs in
|
||||
dependency order). -/
|
||||
def mkTag (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let nId := mkIdent ``Spa.RawId
|
||||
tds.mapM fun td => do
|
||||
let mut pats : Array Term := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map (fun i => mkIdent (.mkSimple s!"a{i}"))
|
||||
let pat ← `($(mkIdent cd.origName) $argNames*)
|
||||
let mut cur : Term ← `(n)
|
||||
let mut lets : Array (Ident × Term) := #[]
|
||||
let mut taggedArgs : Array Term := #[]
|
||||
let mut ri := 0
|
||||
for (f, a) in cd.fields.zip argNames do
|
||||
if f.isRec then
|
||||
let rName := mkIdent (.mkSimple s!"r{ri}")
|
||||
let rhsCall ← `($(mkIdent (tagOf f.recType)) $a $cur)
|
||||
lets := lets.push (rName, rhsCall)
|
||||
taggedArgs := taggedArgs.push (← `($rName |>.1))
|
||||
cur ← `($rName |>.2)
|
||||
ri := ri + 1
|
||||
else
|
||||
taggedArgs := taggedArgs.push a
|
||||
let last := cur
|
||||
let tagged ← `($(mkIdent (taggedOf td.name ++ cd.shortName))
|
||||
(⟨$last⟩ : $nId) $taggedArgs*)
|
||||
let mut body ← `(($tagged, $last + 1))
|
||||
for (rName, rhs) in lets.reverse do
|
||||
body ← `(let $rName := $rhs; $body)
|
||||
pats := pats.push pat
|
||||
rhss := rhss.push body
|
||||
`(command| def $(mkIdent (tagOf td.name)) :
|
||||
$(mkIdent td.name) → Nat → $(mkIdent (taggedOf td.name)) $nId × Nat :=
|
||||
fun e n => match e with $[| $pats => $rhss]*)
|
||||
|
||||
/-- The tag-fold functions: `foldTags f acc t` applies `f` to every tag in `t`,
|
||||
right-to-left, threading `acc`. This is the `Foldable`/`foldr`-over-tags the
|
||||
hand-written collectors (e.g. `subtreeIds`) reduce to. One separate def per
|
||||
family member (the family is a DAG, so no `mutual` block is needed). -/
|
||||
def mkFoldTags (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let τ := mkIdent `τ
|
||||
let m := mkIdent `M
|
||||
let fId := mkIdent `f
|
||||
let accId := mkIdent `acc
|
||||
let tagId := mkIdent `t
|
||||
tds.mapM fun td => do
|
||||
let mut pats : Array Term := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map (fun i => mkIdent (.mkSimple s!"a{i}"))
|
||||
let pat ← `($(mkIdent (taggedOf td.name ++ cd.shortName)) $tagId $argNames*)
|
||||
let mut body : Term := accId
|
||||
for (fld, a) in (cd.fields.zip argNames).reverse do
|
||||
if fld.isRec then
|
||||
body ← `($(mkIdent (foldTagsOf fld.recType)) $fId $body $a)
|
||||
body ← `($fId $tagId $body)
|
||||
pats := pats.push pat
|
||||
rhss := rhss.push body
|
||||
`(command| def $(mkIdent (foldTagsOf td.name)) {$τ:ident : Type} {$m:ident : Type}
|
||||
($fId : $τ → $m → $m) ($accId : $m) :
|
||||
$(mkIdent (taggedOf td.name)) $τ → $m :=
|
||||
fun x => match x with $[| $pats => $rhss]*)
|
||||
|
||||
/-- The well-formedness predicate `T.Tagged.WF : T.Tagged RawId → Prop`: every
|
||||
recursive child's root tag has a strictly smaller postorder index than the node's
|
||||
own tag, and each child is itself well-formed. Leaf constructors are `True`. -/
|
||||
def mkWF (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let tId := mkIdent `t
|
||||
let rawId := mkIdent ``Spa.RawId
|
||||
tds.mapM fun td => do
|
||||
let mut pats : Array Term := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let hasRec := cd.fields.any (·.isRec)
|
||||
let mut patArgs : Array Term := #[]
|
||||
let mut recArgs : Array Ident := #[]
|
||||
let mut i := 0
|
||||
for f in cd.fields do
|
||||
if f.isRec then
|
||||
let a := mkIdent (.mkSimple s!"a{i}")
|
||||
patArgs := patArgs.push a
|
||||
recArgs := recArgs.push a
|
||||
else
|
||||
patArgs := patArgs.push (← `(_))
|
||||
i := i + 1
|
||||
let tagBind : Term ← if hasRec then `($tId) else `(_)
|
||||
let pat ← `($(mkIdent (taggedOf td.name ++ cd.shortName)) $tagBind $patArgs*)
|
||||
let rhs ← if recArgs.isEmpty then `(True) else do
|
||||
let bounds ← recArgs.mapM fun a => `($(a).rootTag.post < $(tId).post)
|
||||
let wfs ← recArgs.mapM fun a => `($(a).WF)
|
||||
mkAndR (bounds ++ wfs)
|
||||
pats := pats.push pat
|
||||
rhss := rhss.push rhs
|
||||
`(command| def $(mkIdent (wfOf td.name)) :
|
||||
$(mkIdent (taggedOf td.name)) $rawId → Prop :=
|
||||
fun x => match x with $[| $pats => $rhss]*)
|
||||
|
||||
/-- The `narrow` coercion `T.Tagged RawId → T.Tagged (Fin N)`, given a bound on
|
||||
the root tag and a well-formedness proof. Each node's tag becomes the `Fin N`
|
||||
built from its postorder index, and recursion threads the bound through `lt_trans`
|
||||
and the (definitionally unfolded) `WF` conjunction. -/
|
||||
def mkNarrow (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let rawId := mkIdent ``Spa.RawId
|
||||
let tId := mkIdent `t
|
||||
let nId := mkIdent `N
|
||||
let hId := mkIdent `h
|
||||
let hwfId := mkIdent `hwf
|
||||
let tgId := mkIdent `tg
|
||||
tds.mapM fun td => do
|
||||
let self ← `($(mkIdent (taggedOf td.name)) $rawId)
|
||||
let mut patss : Array (Array Term) := #[]
|
||||
let mut rhss : Array Term := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map fun i => mkIdent (.mkSimple s!"a{i}")
|
||||
let ctorPat ← `($(mkIdent (taggedOf td.name ++ cd.shortName)) $tgId $argNames*)
|
||||
let k := (cd.fields.filter (·.isRec)).size
|
||||
let mut newArgs : Array Term := #[]
|
||||
let mut ri := 0
|
||||
for (f, a) in cd.fields.zip argNames do
|
||||
if f.isRec then
|
||||
let bound ← projAnd hwfId (ri + 1) (2 * k)
|
||||
let wf ← projAnd hwfId (k + ri + 1) (2 * k)
|
||||
newArgs := newArgs.push (← `($(a).narrow (lt_trans $bound $hId) $wf))
|
||||
ri := ri + 1
|
||||
else
|
||||
newArgs := newArgs.push a
|
||||
let built ← `($(mkIdent (taggedOf td.name ++ cd.shortName)) ⟨$(tgId).post, $hId⟩ $newArgs*)
|
||||
let nPat ← `(_)
|
||||
let hPat ← `($hId)
|
||||
let hwfPat : Term ← if k == 0 then `(_) else `($hwfId)
|
||||
patss := patss.push #[ctorPat, nPat, hPat, hwfPat]
|
||||
rhss := rhss.push built
|
||||
`(command| def $(mkIdent (narrowOf td.name)) : ($tId : $self) → {$nId : ℕ} →
|
||||
$(tId).rootTag.post < $nId → $(tId).WF → $(mkIdent (taggedOf td.name)) (Fin $nId)
|
||||
$[| $[$patss],* => $rhss]*)
|
||||
|
||||
/-- `T.tag_rootTag_post`: the root tag of a freshly tagged node is exactly one
|
||||
below the threaded-out counter, i.e. the node itself is numbered last (postorder).
|
||||
A uniform `cases <;> simp` discharges every constructor. -/
|
||||
def mkTagRootTagPost (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let eId := mkIdent `e
|
||||
let nId := mkIdent `n
|
||||
tds.mapM fun td =>
|
||||
`(command| theorem $(mkIdent (tagRootTagPostOf td.name))
|
||||
($eId : $(mkIdent td.name)) ($nId : ℕ) :
|
||||
($(eId).tag $nId).1.rootTag.post + 1 = ($(eId).tag $nId).2 := by
|
||||
cases $eId:ident <;>
|
||||
simp [$(mkIdent (tagOf td.name)):ident, $(mkIdent (rootTagOf td.name)):ident])
|
||||
|
||||
/-- `T.tag_le`: tagging only ever advances the counter (`n ≤ (e.tag n).2`).
|
||||
Proved by induction; each arm threads the counter through its recursive children
|
||||
(using the relevant `tag_le`/induction hypothesis) and closes with `omega`. -/
|
||||
def mkTagLe (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let eId := mkIdent `e
|
||||
let nId := mkIdent `n
|
||||
tds.mapM fun td => do
|
||||
let mut ctorLabels : Array Ident := #[]
|
||||
let mut binderss : Array (Array Ident) := #[]
|
||||
let mut tacs : Array (TSyntax ``Lean.Parser.Tactic.tacticSeq) := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map fun i => mkIdent (.mkSimple s!"a{i}")
|
||||
let mut ihBinders : Array Ident := #[]
|
||||
let mut haveTacs : Array (TSyntax `tactic) := #[]
|
||||
let mut cur : Term ← `($nId)
|
||||
let mut i := 0
|
||||
for (f, a) in cd.fields.zip argNames do
|
||||
if f.isRec then
|
||||
let fact ← if f.recType == td.name then
|
||||
`($(mkIdent (.mkSimple s!"ih{i}")) $cur)
|
||||
else
|
||||
`($(mkIdent (tagLeOf f.recType)) $a $cur)
|
||||
if f.recType == td.name then
|
||||
ihBinders := ihBinders.push (mkIdent (.mkSimple s!"ih{i}"))
|
||||
haveTacs := haveTacs.push (← `(tactic| have := $fact))
|
||||
cur ← `(($(mkIdent (tagOf f.recType)) $a $cur) |>.2)
|
||||
i := i + 1
|
||||
let simpTac ← `(tactic| simp only [$(mkIdent (tagOf td.name)):ident])
|
||||
let omegaTac ← `(tactic| omega)
|
||||
let allTacs := #[simpTac] ++ haveTacs ++ #[omegaTac]
|
||||
ctorLabels := ctorLabels.push (mkIdent cd.shortName)
|
||||
binderss := binderss.push (argNames ++ ihBinders)
|
||||
tacs := tacs.push (← `(tacticSeq| $[$allTacs]*))
|
||||
`(command| theorem $(mkIdent (tagLeOf td.name)) ($eId : $(mkIdent td.name)) ($nId : ℕ) :
|
||||
$nId ≤ ($(eId).tag $nId).2 := by
|
||||
induction $eId:ident generalizing $nId:ident with
|
||||
$[| $ctorLabels:ident $binderss* => $tacs]*)
|
||||
|
||||
/-- `T.tag_wf`: a freshly tagged term is well-formed. Each recursive child's
|
||||
bound conjunct is closed by `omega` from that child's `tag_rootTag_post` plus the
|
||||
`tag_le` of every later child (which bounds the threaded-out counter), and each
|
||||
well-formedness conjunct is the child's induction hypothesis / `tag_wf`. -/
|
||||
def mkTagWf (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let eId := mkIdent `e
|
||||
let nId := mkIdent `n
|
||||
tds.mapM fun td => do
|
||||
let mut ctorLabels : Array Ident := #[]
|
||||
let mut binderss : Array (Array Ident) := #[]
|
||||
let mut tacs : Array (TSyntax ``Lean.Parser.Tactic.tacticSeq) := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map fun i => mkIdent (.mkSimple s!"a{i}")
|
||||
-- recursive children: (arg, recType, startCounter, sameType?, fieldIndex)
|
||||
let mut recs : Array (Ident × Name × Term × Bool × Nat) := #[]
|
||||
let mut cur : Term ← `($nId)
|
||||
let mut i := 0
|
||||
for (f, a) in cd.fields.zip argNames do
|
||||
if f.isRec then
|
||||
recs := recs.push (a, f.recType, cur, f.recType == td.name, i)
|
||||
cur ← `(($(mkIdent (tagOf f.recType)) $a $cur) |>.2)
|
||||
i := i + 1
|
||||
let k := recs.size
|
||||
let ihBinders := (recs.filter (·.2.2.2.1)).map fun r => mkIdent (.mkSimple s!"ih{r.2.2.2.2}")
|
||||
let tac : TSyntax ``Lean.Parser.Tactic.tacticSeq ← if k == 0 then
|
||||
`(tacticSeq| exact True.intro)
|
||||
else do
|
||||
let mut comps : Array Term := #[]
|
||||
-- bound conjuncts
|
||||
for idx in [0:k] do
|
||||
let (a, rt, s, _, _) := recs[idx]!
|
||||
let mut bHaves : Array (TSyntax `tactic) :=
|
||||
#[← `(tactic| have := $(mkIdent (tagRootTagPostOf rt)) $a $s)]
|
||||
for j in [idx+1:k] do
|
||||
let (aj, rtj, sj, _, _) := recs[j]!
|
||||
bHaves := bHaves.push (← `(tactic| have := $(mkIdent (tagLeOf rtj)) $aj $sj))
|
||||
bHaves := bHaves.push (← `(tactic| omega))
|
||||
comps := comps.push (← `(by $(← `(tacticSeq| $[$bHaves]*))))
|
||||
-- well-formedness conjuncts
|
||||
for idx in [0:k] do
|
||||
let (a, rt, s, same, fi) := recs[idx]!
|
||||
comps := comps.push <| ← if same then `($(mkIdent (.mkSimple s!"ih{fi}")) $s)
|
||||
else `($(mkIdent (tagWfOf rt)) $a $s)
|
||||
let simpTac ← `(tactic| simp only
|
||||
[$(mkIdent (tagOf td.name)):ident, $(mkIdent (wfOf td.name)):ident])
|
||||
let exactTac ← `(tactic| exact ⟨$comps,*⟩)
|
||||
`(tacticSeq| $[$(#[simpTac, exactTac])]*)
|
||||
ctorLabels := ctorLabels.push (mkIdent cd.shortName)
|
||||
binderss := binderss.push (argNames ++ ihBinders)
|
||||
tacs := tacs.push tac
|
||||
`(command| theorem $(mkIdent (tagWfOf td.name)) ($eId : $(mkIdent td.name)) ($nId : ℕ) :
|
||||
($(eId).tag $nId).1.WF := by
|
||||
induction $eId:ident generalizing $nId:ident with
|
||||
$[| $ctorLabels:ident $binderss* => $tacs]*)
|
||||
|
||||
/-- `T.Tagged.narrow_erase`: narrowing the tag type does not change the erased
|
||||
(untagged) term. A per-constructor `simp` with the local `narrow`/`erase`
|
||||
equations, the lower members' `narrow_erase`, and the induction hypotheses. -/
|
||||
def mkNarrowErase (tds : Array TypeData) : CommandElabM (Array (TSyntax `command)) := do
|
||||
let rawId := mkIdent ``Spa.RawId
|
||||
let tId := mkIdent `t
|
||||
let nId := mkIdent `N
|
||||
let hId := mkIdent `h
|
||||
let hwfId := mkIdent `hwf
|
||||
let tgId := mkIdent `tg
|
||||
tds.mapM fun td => do
|
||||
let mut ctorLabels : Array Ident := #[]
|
||||
let mut binderss : Array (Array Ident) := #[]
|
||||
let mut tacs : Array (TSyntax ``Lean.Parser.Tactic.tacticSeq) := #[]
|
||||
for cd in td.ctors do
|
||||
let argNames := (Array.range cd.fields.size).map fun i => mkIdent (.mkSimple s!"a{i}")
|
||||
let mut lemmas : Array Term :=
|
||||
#[← `($(mkIdent (narrowOf td.name))), ← `($(mkIdent (eraseOf td.name)))]
|
||||
let mut ihBinders : Array Ident := #[]
|
||||
let mut seenLower : Array Name := #[]
|
||||
let mut i := 0
|
||||
for f in cd.fields do
|
||||
if f.isRec then
|
||||
if f.recType == td.name then
|
||||
let ih := mkIdent (.mkSimple s!"ih{i}")
|
||||
ihBinders := ihBinders.push ih
|
||||
lemmas := lemmas.push (← `($ih))
|
||||
else if !seenLower.contains f.recType then
|
||||
seenLower := seenLower.push f.recType
|
||||
lemmas := lemmas.push (← `($(mkIdent (narrowEraseOf f.recType))))
|
||||
i := i + 1
|
||||
let introTac ← `(tactic| intro $nId $hId $hwfId)
|
||||
let simpTac ← `(tactic| simp [$[$lemmas:term],*])
|
||||
ctorLabels := ctorLabels.push (mkIdent cd.shortName)
|
||||
binderss := binderss.push (#[tgId] ++ argNames ++ ihBinders)
|
||||
tacs := tacs.push (← `(tacticSeq| $[$(#[introTac, simpTac])]*))
|
||||
`(command| theorem $(mkIdent (narrowEraseOf td.name)) :
|
||||
($tId : $(mkIdent (taggedOf td.name)) $rawId) → ∀ {$nId : ℕ}
|
||||
($hId : $(tId).rootTag.post < $nId) ($hwfId : $(tId).WF),
|
||||
($(tId).narrow $hId $hwfId).erase = $(tId).erase := by
|
||||
intro $tId:ident
|
||||
induction $tId:ident with
|
||||
$[| $ctorLabels:ident $binderss* => $tacs]*)
|
||||
|
||||
/-- `derive_tagged T₁ … Tₙ` — generate tagged mirrors, `erase`, and `tag` for the
|
||||
given family of inductives. -/
|
||||
syntax (name := deriveTaggedCmd) "derive_tagged " ident+ : command
|
||||
|
||||
@[command_elab deriveTaggedCmd]
|
||||
def elabDeriveTagged : CommandElab := fun stx => do
|
||||
match stx with
|
||||
| `(derive_tagged $ids*) =>
|
||||
let family ← ids.mapM fun i => Command.liftCoreM (realizeGlobalConstNoOverload i)
|
||||
let τ := mkIdent `τ
|
||||
let tds ← Command.liftTermElabM (gather family τ)
|
||||
for d in (← mkInductives tds τ) do elabCommand d
|
||||
for d in (← mkDeriveInstances tds) do elabCommand d
|
||||
for d in (← mkRootTag tds) do elabCommand d
|
||||
for d in (← mkErase tds) do elabCommand d
|
||||
for d in (← mkTag tds) do elabCommand d
|
||||
for d in (← mkFoldTags tds) do elabCommand d
|
||||
for d in (← mkWF tds) do elabCommand d
|
||||
for d in (← mkNarrow tds) do elabCommand d
|
||||
for d in (← mkTagRootTagPost tds) do elabCommand d
|
||||
for d in (← mkTagLe tds) do elabCommand d
|
||||
for d in (← mkTagWf tds) do elabCommand d
|
||||
for d in (← mkNarrowErase tds) do elabCommand d
|
||||
| _ => throwUnsupportedSyntax
|
||||
|
||||
end Spa.DeriveTagged
|
||||
104
lean/Spa/Language/Tagged/Graphs.lean
Normal file
104
lean/Spa/Language/Tagged/Graphs.lean
Normal file
@@ -0,0 +1,104 @@
|
||||
import Spa.Language
|
||||
import Spa.Language.Graphs
|
||||
import Spa.Language.Tagged.Basic
|
||||
import Spa.Language.Tagged.Properties
|
||||
|
||||
namespace Spa
|
||||
|
||||
open GGraph
|
||||
|
||||
def Stmt.Tagged.cfg {τ : Type} : Stmt.Tagged τ → GGraph (Option (BasicStmt.Tagged τ))
|
||||
| .basic _ bs => GGraph.singleton (some bs)
|
||||
| .andThen _ s₁ s₂ => s₁.cfg ⤳ s₂.cfg
|
||||
| .ifElse _ _ s₁ s₂ => s₁.cfg ∙ s₂.cfg
|
||||
| .whileLoop _ _ s => GGraph.loop s.cfg
|
||||
|
||||
theorem Stmt.Tagged.cfg_graph {τ : Type} : ∀ (t : Stmt.Tagged τ),
|
||||
(Option.map BasicStmt.Tagged.erase) <$> t.cfg = t.erase.cfg
|
||||
| .basic _ bs => by simp [Stmt.Tagged.cfg, Stmt.cfg, Stmt.Tagged.erase, BasicStmt.Tagged.erase]
|
||||
| .andThen _ s₁ s₂ => by
|
||||
simp [Stmt.Tagged.cfg, Stmt.cfg, Stmt.Tagged.erase, Stmt.Tagged.cfg_graph s₁, Stmt.Tagged.cfg_graph s₂]
|
||||
| .ifElse _ _ s₁ s₂ => by
|
||||
simp [Stmt.Tagged.cfg, Stmt.cfg, Stmt.Tagged.erase, Stmt.Tagged.cfg_graph s₁, Stmt.Tagged.cfg_graph s₂]
|
||||
| .whileLoop _ _ s => by
|
||||
simp [Stmt.Tagged.cfg, Stmt.cfg, Stmt.Tagged.erase, Stmt.Tagged.cfg_graph s]
|
||||
|
||||
def GGraph.nodeLabel {τ : Type} (g : GGraph (Option (BasicStmt.Tagged τ))) (i : g.Index) :
|
||||
Option τ :=
|
||||
(g.nodes i).map BasicStmt.Tagged.rootTag
|
||||
|
||||
def GGraph.stateOf {τ : Type} [DecidableEq τ] (g : GGraph (Option (BasicStmt.Tagged τ)))
|
||||
(id : τ) : Option g.Index :=
|
||||
g.indices.find? (fun i => decide (g.nodeLabel i = some id))
|
||||
|
||||
theorem GGraph.stateOf_label {τ : Type} [DecidableEq τ]
|
||||
{g : GGraph (Option (BasicStmt.Tagged τ))} {id : τ}
|
||||
{i : g.Index} (h : g.stateOf id = some i) : g.nodeLabel i = some id := by
|
||||
rw [GGraph.stateOf] at h
|
||||
simpa using List.find?_some h
|
||||
|
||||
namespace Program
|
||||
|
||||
variable (p : Program)
|
||||
|
||||
def tagged : Stmt.Tagged RawId := tagStmt p.rootStmt
|
||||
|
||||
def size : ℕ := p.tagged.rootTag.post + 1
|
||||
|
||||
theorem size_pos : 0 < p.size := Nat.succ_pos _
|
||||
|
||||
abbrev NodeId : Type := Fin p.size
|
||||
|
||||
theorem tagged_wf : p.tagged.WF := Stmt.tag_wf p.rootStmt 0
|
||||
|
||||
def taggedFin : Stmt.Tagged p.NodeId :=
|
||||
p.tagged.narrow (Nat.lt_succ_self _) p.tagged_wf
|
||||
|
||||
def taggedCfg : GGraph (Option (BasicStmt.Tagged p.NodeId)) :=
|
||||
GGraph.wrap p.taggedFin.cfg
|
||||
|
||||
theorem taggedCfg_erase :
|
||||
(Option.map BasicStmt.Tagged.erase) <$> p.taggedCfg = p.cfg := by
|
||||
rw [taggedCfg, GGraph.map_wrap, Stmt.Tagged.cfg_graph, taggedFin,
|
||||
Stmt.Tagged.narrow_erase, tagged, erase_tagStmt]
|
||||
rfl
|
||||
|
||||
theorem taggedCfg_size : p.taggedCfg.size = p.cfg.size := by
|
||||
conv_rhs => rw [← p.taggedCfg_erase]
|
||||
rfl
|
||||
|
||||
def nodeIdOf (s : p.State) : Option p.NodeId :=
|
||||
p.taggedCfg.nodeLabel (Fin.cast p.taggedCfg_size.symm s)
|
||||
|
||||
def stateOfNodeId (id : p.NodeId) : Option p.State :=
|
||||
(p.taggedCfg.stateOf id).map (Fin.cast p.taggedCfg_size)
|
||||
|
||||
theorem cfg_nodes_eq (s : p.State) :
|
||||
p.cfg.nodes s = Option.map BasicStmt.Tagged.erase
|
||||
(p.taggedCfg.nodes (Fin.cast p.taggedCfg_size.symm s)) := by
|
||||
have key : ∀ (g : Graph) (hsz : p.taggedCfg.size = g.size),
|
||||
(Option.map BasicStmt.Tagged.erase) <$> p.taggedCfg = g →
|
||||
∀ i : Fin g.size,
|
||||
g.nodes i = Option.map BasicStmt.Tagged.erase
|
||||
(p.taggedCfg.nodes (Fin.cast hsz.symm i)) := by
|
||||
intro g hsz hg i
|
||||
subst hg
|
||||
rfl
|
||||
exact key p.cfg p.taggedCfg_size p.taggedCfg_erase s
|
||||
|
||||
theorem nodeIdOf_isSome_of_code {s : p.State} {bs : BasicStmt}
|
||||
(h : p.code s = some bs) : (p.nodeIdOf s).isSome = true := by
|
||||
have hc : Option.map BasicStmt.Tagged.erase
|
||||
(p.taggedCfg.nodes (Fin.cast p.taggedCfg_size.symm s)) = some bs := by
|
||||
rw [← p.cfg_nodes_eq s]; exact h
|
||||
unfold Program.nodeIdOf GGraph.nodeLabel
|
||||
cases hcase : p.taggedCfg.nodes (Fin.cast p.taggedCfg_size.symm s) with
|
||||
| none => rw [hcase] at hc; simp at hc
|
||||
| some tbs => simp
|
||||
|
||||
def nodeIdOfNonempty (s : p.State) {bs : BasicStmt} (h : p.code s = some bs) : p.NodeId :=
|
||||
(p.nodeIdOf s).get (p.nodeIdOf_isSome_of_code h)
|
||||
|
||||
end Program
|
||||
|
||||
end Spa
|
||||
9
lean/Spa/Language/Tagged/Id.lean
Normal file
9
lean/Spa/Language/Tagged/Id.lean
Normal file
@@ -0,0 +1,9 @@
|
||||
import Mathlib.Data.Nat.Notation
|
||||
|
||||
namespace Spa
|
||||
|
||||
structure RawId where
|
||||
post : ℕ
|
||||
deriving DecidableEq, Repr
|
||||
|
||||
end Spa
|
||||
29
lean/Spa/Language/Tagged/Properties.lean
Normal file
29
lean/Spa/Language/Tagged/Properties.lean
Normal file
@@ -0,0 +1,29 @@
|
||||
import Spa.Language.Tagged.Basic
|
||||
|
||||
namespace Spa
|
||||
|
||||
@[simp] theorem Expr.erase_tag (e : Expr) (n : ℕ) : (e.tag n).1.erase = e := by
|
||||
induction e generalizing n with
|
||||
| add a b iha ihb => simp [Expr.tag, Expr.Tagged.erase, iha, ihb]
|
||||
| sub a b iha ihb => simp [Expr.tag, Expr.Tagged.erase, iha, ihb]
|
||||
| var x => simp [Expr.tag, Expr.Tagged.erase]
|
||||
| num k => simp [Expr.tag, Expr.Tagged.erase]
|
||||
|
||||
@[simp] theorem BasicStmt.erase_tag (bs : BasicStmt) (n : ℕ) :
|
||||
(bs.tag n).1.erase = bs := by
|
||||
cases bs with
|
||||
| assign x e => simp [BasicStmt.tag, BasicStmt.Tagged.erase]
|
||||
| noop => simp [BasicStmt.tag, BasicStmt.Tagged.erase]
|
||||
|
||||
@[simp] theorem Stmt.erase_tag (s : Stmt) (n : ℕ) : (s.tag n).1.erase = s := by
|
||||
induction s generalizing n with
|
||||
| basic bs => simp [Stmt.tag, Stmt.Tagged.erase]
|
||||
| andThen a b iha ihb => simp [Stmt.tag, Stmt.Tagged.erase, iha, ihb]
|
||||
| ifElse e a b iha ihb => simp [Stmt.tag, Stmt.Tagged.erase, iha, ihb]
|
||||
| whileLoop e s ih => simp [Stmt.tag, Stmt.Tagged.erase, ih]
|
||||
|
||||
/-- Erasing a freshly tagged program recovers it. -/
|
||||
theorem erase_tagStmt (s : Stmt) : (tagStmt s).erase = s := by
|
||||
simp [tagStmt]
|
||||
|
||||
end Spa
|
||||
46
lean/Spa/Language/Tagged/TODO.md
Normal file
46
lean/Spa/Language/Tagged/TODO.md
Normal file
@@ -0,0 +1,46 @@
|
||||
# Tagged AST — follow-ups
|
||||
|
||||
## Descendant tracking — parked
|
||||
|
||||
The interval-labeling descendant test and its correctness proof
|
||||
(`descendant_iff_tagStmt` and supporting rose-tree/`Good` machinery) have been
|
||||
removed from the live code and parked in `DESCENDANT-TRACKING.md`, with a revival
|
||||
checklist. It's a computational optimization not yet needed; revive it (and the
|
||||
`NodeId.desc` field) when LICM wants fast ancestor queries.
|
||||
|
||||
## ID → CFG-state mapping — plan part B — DONE
|
||||
|
||||
`Graphs.lean` now defines a payload-generic `GGraph α` (with `Graph := GGraph
|
||||
(List BasicStmt)` as the concrete CFG), so the labelled CFG **reuses** the graph
|
||||
combinators instead of mirroring them. In `Cfg.lean`:
|
||||
`buildCfgL : Stmt.Tagged NodeId → GGraph (List (BasicStmt.Tagged NodeId))` is just
|
||||
`buildCfg` at the tagged payload; `buildCfgL_graph :
|
||||
(buildCfgL t).map (List.map erase) = buildCfg t.erase` connects it to the real
|
||||
CFG; and `GGraph.nodeLabel`/`GGraph.stateOf` read a node's id straight from its
|
||||
payload (`stateOf_label` is the soundness). No `LGraph`, no separate `label`
|
||||
field, no duplicated combinators.
|
||||
|
||||
## ID → CFG-state mapping — totality — DONE
|
||||
|
||||
The `Option`-valued `nodeIdOf`/`stateOfNodeId` are now proven total on the inputs
|
||||
that matter (`Graphs.lean`), via a payload-list characterization of the CFG:
|
||||
|
||||
- `GGraph.nodeList` flattens `nodes` into the list of payloads, with combinator
|
||||
lemmas (`nodeList_comp/link/loop/wrap`) reducing it through the CFG builders.
|
||||
- `Stmt.Tagged.basics` lists a program's basic statements; the master lemma
|
||||
`Stmt.Tagged.cfg_nodeList_filter` (and its program-level
|
||||
`taggedCfg_nodeList_filter`) shows the non-empty CFG nodes are *exactly* the
|
||||
singletons `[bs]` for `bs ∈ basics`.
|
||||
- AST ⇒ CFG: `exists_state_of_mem_basics` (a state with payload `[bs]`) and
|
||||
`stateOfNodeId_isSome` (the search succeeds).
|
||||
- CFG ⇒ AST: `exists_basic_of_code_ne_nil` (a non-empty node is `[bs]`, with
|
||||
`code = [bs.erase]` and `nodeIdOf = some bs.rootTag`) and `nodeIdOf_isSome`.
|
||||
|
||||
All `propext`/`Quot.sound`-only (no `sorry`, no choice).
|
||||
|
||||
Remaining nice-to-have:
|
||||
- Injectivity: distinct basic-statement ids map to distinct states, giving a
|
||||
two-sided id ↔ state correspondence (upgrading the existence results above to a
|
||||
genuine bijection, and pinning `stateOfNodeId (bs.rootTag)` to *the* state
|
||||
holding `bs`). The `tag`-uniqueness fact this needs (`Nodup` of postorder tags)
|
||||
was part of the parked descendant machinery in `DESCENDANT-TRACKING.md`.
|
||||
@@ -1,16 +1,46 @@
|
||||
import Spa.Language.Semantics
|
||||
import Spa.Language.Graphs
|
||||
|
||||
/-!
|
||||
|
||||
# Program Traces
|
||||
|
||||
This module defines program traces tied to Control Flow Graphs, or CFGs
|
||||
(see `Spa.GGraph` and `Spa.Graph`). These traces boil town to sequences of
|
||||
basic-block executions (really, `Spa.BasicStmt` executions), each of which must
|
||||
have an actual basic block in the graph _and_ be connected to the previous
|
||||
basic block by an edge. In this way, traces encode executions admitted
|
||||
by the CFG.
|
||||
|
||||
While the regular `Trace` is just _any_ path through the graph, an
|
||||
`EndToEndTrace` is a path from the entry node to the exit node, denoting
|
||||
full program execution.
|
||||
|
||||
Properties about graphs and language semantics (especially,
|
||||
the fact that the graph contains the proper basic block and edges
|
||||
to represent any program execution according to the
|
||||
language's big-step semantics `EvalStmt`) is found
|
||||
in `Spa/Language/Properties.lean`.
|
||||
|
||||
-/
|
||||
|
||||
namespace Spa
|
||||
|
||||
inductive Trace (g : Graph) : g.Index → g.Index → Env → Env → Prop
|
||||
/-- A partial trace through a graph `g`, starting right before
|
||||
the execution of the basic block at the first index, and
|
||||
ending right after the execution of the basic block at the last index. -/
|
||||
inductive Trace (g : Graph) : g.Index → g.Index → Env → Env → Type
|
||||
| single {ρ₁ ρ₂ : Env} {idx : g.Index} :
|
||||
EvalBasicStmts ρ₁ (g.nodes idx) ρ₂ → Trace g idx idx ρ₁ ρ₂
|
||||
EvalBasicStmtOpt ρ₁ (g.nodes idx) ρ₂ → Trace g idx idx ρ₁ ρ₂
|
||||
| edge {ρ₁ ρ₂ ρ₃ : Env} {idx₁ idx₂ idx₃ : g.Index} :
|
||||
EvalBasicStmts ρ₁ (g.nodes idx₁) ρ₂ → (idx₁, idx₂) ∈ g.edges →
|
||||
EvalBasicStmtOpt ρ₁ (g.nodes idx₁) ρ₂ → (idx₁, idx₂) ∈ g.edges →
|
||||
Trace g idx₂ idx₃ ρ₂ ρ₃ → Trace g idx₁ idx₃ ρ₁ ρ₃
|
||||
|
||||
lemma Trace.concat {g : Graph} {idx₁ idx₂ idx₃ idx₄ : g.Index}
|
||||
/-- Sequence two traces together. Since the endpoint of the first trace
|
||||
is _after_ its last basic block's execution, and the beginning of
|
||||
the next trace is _before_ its first basic block's execution,
|
||||
there must be an edge to connect the two. -/
|
||||
noncomputable def Trace.concat {g : Graph} {idx₁ idx₂ idx₃ idx₄ : g.Index}
|
||||
{ρ₁ ρ₂ ρ₃ : Env} (tr₁ : Trace g idx₁ idx₂ ρ₁ ρ₂)
|
||||
(he : (idx₂, idx₃) ∈ g.edges) (tr₂ : Trace g idx₃ idx₄ ρ₂ ρ₃) :
|
||||
Trace g idx₁ idx₄ ρ₁ ρ₃ := by
|
||||
@@ -18,7 +48,10 @@ lemma Trace.concat {g : Graph} {idx₁ idx₂ idx₃ idx₄ : g.Index}
|
||||
| single hbs => exact Trace.edge hbs he tr₂
|
||||
| edge hbs he' _ ih => exact Trace.edge hbs he' (ih he tr₂)
|
||||
|
||||
inductive EndToEndTrace (g : Graph) (ρ₁ ρ₂ : Env) : Prop
|
||||
scoped notation:65 tr₁:66 " ++< " he " >++ " tr₂:65 => Trace.concat tr₁ he tr₂
|
||||
|
||||
/-- A beginning-to-end trace corresponding to the CFG `g`. -/
|
||||
inductive EndToEndTrace (g : Graph) (ρ₁ ρ₂ : Env) : Type
|
||||
| intro (idx₁ : g.Index) (idx₁_mem : idx₁ ∈ g.inputs)
|
||||
(idx₂ : g.Index) (idx₂_mem : idx₂ ∈ g.outputs)
|
||||
(trace : Trace g idx₁ idx₂ ρ₁ ρ₂) : EndToEndTrace g ρ₁ ρ₂
|
||||
|
||||
@@ -12,6 +12,18 @@ etc.. What remains are a couple of theorems about folds, as well
|
||||
as `FiniteHeightLattice`, the core concept of lattice-based static
|
||||
program analyses. See the documentation on that class for more information. -/
|
||||
|
||||
namespace Option
|
||||
|
||||
/-- Equality-sensitive eliminator for options in which the `some` case
|
||||
is sensitive to the base `β`. This makes it mirror a one-element fold
|
||||
more closely. -/
|
||||
def elimEq {α : Type*} {β : Sort*} :
|
||||
(o : Option α) → β → ((a : α) → o = some a → β → β) → β
|
||||
| none, b, _ => b
|
||||
| some a, b, f => f a rfl b
|
||||
|
||||
end Option
|
||||
|
||||
namespace Spa
|
||||
|
||||
/-- Predicate for binary functions independently monotone in both their arguments. -/
|
||||
@@ -61,6 +73,16 @@ lemma foldl_mono' (l : List α) (f : β → α → β)
|
||||
| nil => exact hb
|
||||
| cons x xs ih => exact ih (hf x hb)
|
||||
|
||||
omit [Preorder α] in
|
||||
/-- The equality-aware eliminator (that also alters its behavior dependent on base case)
|
||||
for option is monotonic. -/
|
||||
lemma elimEq_self_mono (o : Option α) (g : (a : α) → o = some a → β → β)
|
||||
(hg : ∀ a h, Monotone (g a h)) :
|
||||
Monotone (o.elimEq · g) := by
|
||||
cases o with
|
||||
| none => exact monotone_id
|
||||
| some a => exact hg a rfl
|
||||
|
||||
end Folds
|
||||
|
||||
/-- Predicate on types with `Preorder` that claims all $<$ chains in the type have at most `n` comparisons. -/
|
||||
@@ -76,69 +98,53 @@ lemma boundedChains_of_subsingleton (α : Type*) [Preorder α] [Subsingleton α]
|
||||
exact (c.step ⟨0, by omega⟩).ne (Subsingleton.elim _ _)
|
||||
|
||||
/-- A finite height lattice is a lattice in which all chains $a < \ldots < z$ have a maximum height `height`. -/
|
||||
class FiniteHeightLattice (α : Type*) extends Lattice α where
|
||||
longestChain : LTSeries α
|
||||
chains_bounded : BoundedChains α longestChain.length
|
||||
class FiniteHeightLattice (α : Type*) extends Lattice α, OrderBot α, OrderTop α where
|
||||
height : ℕ
|
||||
chains_bounded : BoundedChains α height
|
||||
|
||||
-- a < ... < z
|
||||
-- ----------- length <= height
|
||||
|
||||
namespace FiniteHeightLattice
|
||||
|
||||
def height (α : Type*) [FiniteHeightLattice α] : ℕ :=
|
||||
(longestChain (α := α)).length
|
||||
|
||||
variable (α : Type*) [FiniteHeightLattice α]
|
||||
|
||||
instance (priority := 100) : Bot α := ⟨(longestChain (α := α)).head⟩
|
||||
instance (priority := 100) : Top α := ⟨(longestChain (α := α)).last⟩
|
||||
|
||||
/-- The bottom element `⊥` of a finite height lattice is _actually_ the least element. -/
|
||||
lemma bot_le (a : α) : (⊥ : α) ≤ a := by
|
||||
by_cases heq : ⊥ ⊓ a = ⊥
|
||||
· exact inf_eq_left.mp heq
|
||||
· exfalso
|
||||
have hlt : ⊥ ⊓ a < (longestChain (α := α)).head :=
|
||||
lt_of_le_of_ne inf_le_left heq
|
||||
have hbound := chains_bounded ((longestChain (α := α)).cons (⊥ ⊓ a) hlt)
|
||||
rw [RelSeries.cons_length] at hbound
|
||||
omega
|
||||
|
||||
/-- The top element `⊤` of a finite height lattice is _actually_ the greatest element. -/
|
||||
lemma le_top (a : α) : a ≤ (⊤ : α) := by
|
||||
by_cases heq : a ⊔ ⊤ = ⊤
|
||||
· exact sup_eq_right.mp heq
|
||||
· exfalso
|
||||
have hlt : (longestChain (α := α)).last < a ⊔ ⊤ :=
|
||||
lt_of_le_of_ne le_sup_right (Ne.symm heq)
|
||||
have hbound := chains_bounded ((longestChain (α := α)).snoc (a ⊔ ⊤) hlt)
|
||||
rw [RelSeries.snoc_length] at hbound
|
||||
omega
|
||||
|
||||
/-- This is something like a lemma about isomorphic types having the same height.
|
||||
Given a finite-height lattice `α`, lattice `β`, and a `Monotone` bijection
|
||||
between the two, we can show that lattice `β` also has a finite height.
|
||||
|
||||
The proof is fairly trivial: the longest chain in `α` can be transported
|
||||
to be a longest chain in `β` (by monotonicity), establishing a height witness.
|
||||
At the same time, any chain in `β` can be transported to a chain in `α`,
|
||||
The proof is fairly trivial: any chain in `β` can be transported to a chain in `α`,
|
||||
and must be bounded by the same height by `FiniteHeightLattice.chains_bounded`. -/
|
||||
def transport {α β : Type*} [Lattice β]
|
||||
[I : FiniteHeightLattice α] (f : α → β) (g : β → α)
|
||||
(hf : Monotone f) (hg : Monotone g)
|
||||
(hgf : Function.LeftInverse g f) (hfg : Function.LeftInverse f g) :
|
||||
(hfg : Function.LeftInverse f g) :
|
||||
FiniteHeightLattice β where
|
||||
toLattice := inferInstance
|
||||
longestChain :=
|
||||
I.longestChain.map f (hf.strictMono_of_injective hgf.injective)
|
||||
toOrderBot := {
|
||||
bot := f (⊥ : α)
|
||||
bot_le := fun b => by
|
||||
rw [← hfg b]
|
||||
exact hf (_root_.bot_le : (⊥ : α) ≤ g b) }
|
||||
toOrderTop := {
|
||||
top := f (⊤ : α)
|
||||
le_top := fun b => by
|
||||
rw [← hfg b]
|
||||
exact hf (_root_.le_top : g b ≤ (⊤ : α)) }
|
||||
height := I.height
|
||||
chains_bounded := fun c =>
|
||||
I.chains_bounded (c.map g (hg.strictMono_of_injective hfg.injective))
|
||||
|
||||
/-- A `Unique` lattice trivially has finite height: its only chain is the singleton
|
||||
`[default]`, and there are no nontrivial `<` chains in a subsingleton. -/
|
||||
def ofUnique (α : Type*) [Lattice α] [Unique α] : FiniteHeightLattice α where
|
||||
def ofUnique (α : Type*) [Lattice α] [Unique α] :
|
||||
FiniteHeightLattice α where
|
||||
toLattice := inferInstance
|
||||
longestChain := RelSeries.singleton _ default
|
||||
toOrderBot := {
|
||||
bot := default
|
||||
bot_le := fun _ => le_of_eq (Subsingleton.elim _ _) }
|
||||
toOrderTop := {
|
||||
top := default
|
||||
le_top := fun _ => le_of_eq (Subsingleton.elim _ _) }
|
||||
height := 0
|
||||
chains_bounded := boundedChains_of_subsingleton α 0
|
||||
|
||||
end FiniteHeightLattice
|
||||
|
||||
@@ -10,6 +10,8 @@ inductive AboveBelow (α : Type*) where
|
||||
|
||||
namespace AboveBelow
|
||||
|
||||
attribute [aesop safe cases] AboveBelow
|
||||
|
||||
instance {α : Type*} [ToString α] : ToString (AboveBelow α) where
|
||||
toString
|
||||
| bot => "⊥"
|
||||
@@ -49,36 +51,22 @@ instance : Min (AboveBelow α) where
|
||||
(mk x ⊓ mk y : AboveBelow α) = if x = y then mk x else bot := rfl
|
||||
|
||||
protected lemma sup_comm (a b : AboveBelow α) : a ⊔ b = b ⊔ a := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;> simp only
|
||||
[bot_sup, sup_bot, top_sup, sup_top, mk_sup_mk]
|
||||
split_ifs with h₁ h₂ h₂ <;> simp_all
|
||||
aesop
|
||||
|
||||
protected lemma sup_assoc (a b c : AboveBelow α) : a ⊔ b ⊔ c = a ⊔ (b ⊔ c) := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;> rcases c with _ | _ | z <;>
|
||||
simp only [bot_sup, sup_bot, top_sup, sup_top, mk_sup_mk]
|
||||
split_ifs <;> simp_all
|
||||
aesop
|
||||
|
||||
protected lemma inf_comm (a b : AboveBelow α) : a ⊓ b = b ⊓ a := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;> simp only
|
||||
[bot_inf, inf_bot, top_inf, inf_top, mk_inf_mk]
|
||||
split_ifs with h₁ h₂ h₂ <;> simp_all
|
||||
aesop
|
||||
|
||||
protected lemma inf_assoc (a b c : AboveBelow α) : a ⊓ b ⊓ c = a ⊓ (b ⊓ c) := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;> rcases c with _ | _ | z <;>
|
||||
simp only [bot_inf, inf_bot, top_inf, inf_top, mk_inf_mk]
|
||||
split_ifs <;> simp_all
|
||||
aesop
|
||||
|
||||
protected lemma sup_inf_self (a b : AboveBelow α) : a ⊔ a ⊓ b = a := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;>
|
||||
simp only [bot_sup, sup_bot, top_sup, sup_top, mk_sup_mk,
|
||||
bot_inf, inf_bot, top_inf, inf_top, mk_inf_mk] <;>
|
||||
try (split_ifs <;> simp_all)
|
||||
aesop
|
||||
|
||||
protected lemma inf_sup_self (a b : AboveBelow α) : a ⊓ (a ⊔ b) = a := by
|
||||
rcases a with _ | _ | x <;> rcases b with _ | _ | y <;>
|
||||
simp only [bot_sup, sup_bot, top_sup, sup_top, mk_sup_mk,
|
||||
bot_inf, inf_bot, top_inf, inf_top, mk_inf_mk] <;>
|
||||
try (split_ifs <;> simp_all)
|
||||
aesop
|
||||
|
||||
instance : Lattice (AboveBelow α) :=
|
||||
Lattice.mk' AboveBelow.sup_comm AboveBelow.sup_assoc
|
||||
@@ -93,6 +81,14 @@ lemma bot_le' (a : AboveBelow α) : (bot : AboveBelow α) ≤ a :=
|
||||
lemma le_top' (a : AboveBelow α) : a ≤ (top : AboveBelow α) :=
|
||||
le_iff.mpr (sup_top a)
|
||||
|
||||
instance : OrderBot (AboveBelow α) where
|
||||
bot := bot
|
||||
bot_le := bot_le'
|
||||
|
||||
instance : OrderTop (AboveBelow α) where
|
||||
top := top
|
||||
le_top := le_top'
|
||||
|
||||
lemma bot_lt_mk (x : α) : (bot : AboveBelow α) < mk x :=
|
||||
lt_of_le_of_ne (bot_le' _) (by simp)
|
||||
|
||||
@@ -195,13 +191,7 @@ def rank : AboveBelow α → ℕ
|
||||
lemma not_mk_lt_mk (x y : α) : ¬(mk x : AboveBelow α) < mk y := by
|
||||
intro h
|
||||
obtain ⟨hle, hne⟩ := lt_iff_le_and_ne.mp h
|
||||
have hsup := le_iff.mp hle
|
||||
rw [mk_sup_mk] at hsup
|
||||
by_cases hxy : x = y
|
||||
· rw [if_pos hxy] at hsup
|
||||
exact hne hsup
|
||||
· rw [if_neg hxy] at hsup
|
||||
exact absurd hsup (by simp)
|
||||
rcases le_cases hle with h | h | h <;> simp_all
|
||||
|
||||
lemma rank_strictMono : StrictMono (rank : AboveBelow α → ℕ) := by
|
||||
intro a b hab
|
||||
@@ -224,10 +214,9 @@ lemma boundedChains : BoundedChains (AboveBelow α) 2 := fun c => by
|
||||
|
||||
instance [Inhabited α] : FiniteHeightLattice (AboveBelow α) where
|
||||
toLattice := inferInstance
|
||||
longestChain :=
|
||||
((RelSeries.singleton _ bot).snoc (mk default)
|
||||
(by rw [RelSeries.last_singleton]; exact bot_lt_mk default)).snoc top
|
||||
(by rw [RelSeries.last_snoc]; exact mk_lt_top default)
|
||||
toOrderBot := inferInstance
|
||||
toOrderTop := inferInstance
|
||||
height := 2
|
||||
chains_bounded := boundedChains
|
||||
|
||||
end AboveBelow
|
||||
|
||||
@@ -29,8 +29,9 @@ lemma boundedChains : BoundedChains Bool 1 := fun c => by
|
||||
|
||||
instance : FiniteHeightLattice Bool where
|
||||
toLattice := inferInstance
|
||||
longestChain := (RelSeries.singleton _ (⊥ : Bool)).snoc (⊤ : Bool)
|
||||
(by rw [RelSeries.last_singleton]; exact bot_lt_top)
|
||||
toOrderBot := inferInstance
|
||||
toOrderTop := inferInstance
|
||||
height := 1
|
||||
chains_bounded := boundedChains
|
||||
|
||||
end Bool
|
||||
|
||||
@@ -72,6 +72,12 @@ lemma mem_sup {fm₁ fm₂ : FiniteMap A B ks} {k : A} {v : B}
|
||||
obtain ⟨i, hi, rfl⟩ := h
|
||||
exact ⟨fm₁ i, fm₂ i, rfl, ⟨i, hi, rfl⟩, ⟨i, hi, rfl⟩⟩
|
||||
|
||||
lemma mem_inf {fm₁ fm₂ : FiniteMap A B ks} {k : A} {v : B}
|
||||
(h : (k, v) ∈ fm₁ ⊓ fm₂) :
|
||||
∃ v₁ v₂, v = v₁ ⊓ v₂ ∧ (k, v₁) ∈ fm₁ ∧ (k, v₂) ∈ fm₂ := by
|
||||
obtain ⟨i, hi, rfl⟩ := h
|
||||
exact ⟨fm₁ i, fm₂ i, rfl, ⟨i, hi, rfl⟩, ⟨i, hi, rfl⟩⟩
|
||||
|
||||
section Updating
|
||||
|
||||
variable [DecidableEq A]
|
||||
|
||||
@@ -107,62 +107,18 @@ section FiniteHeight
|
||||
|
||||
variable [FiniteHeightLattice β]
|
||||
|
||||
private lemma consBot_strictMono {n : ℕ} :
|
||||
StrictMono (fun b : β => (Fin.cons b (⊥ : Fin n → β) : Fin (n + 1) → β)) := by
|
||||
intro a b hab
|
||||
refine lt_iff_le_and_ne.mpr ⟨?_, ?_⟩
|
||||
· refine Pi.le_def.mpr (fun i => Fin.cases ?_ (fun j => ?_) i)
|
||||
· simpa using hab.le
|
||||
· simp
|
||||
· exact fun h => hab.ne (by simpa using congrFun h 0)
|
||||
|
||||
private lemma consTop_strictMono {n : ℕ} :
|
||||
StrictMono (fun f : Fin n → β => (Fin.cons (⊤ : β) f : Fin (n + 1) → β)) := by
|
||||
intro f g hfg
|
||||
refine lt_iff_le_and_ne.mpr ⟨?_, ?_⟩
|
||||
· refine Pi.le_def.mpr (fun i => Fin.cases ?_ (fun j => ?_) i)
|
||||
· simp
|
||||
· simpa using Pi.le_def.mp hfg.le j
|
||||
· intro h
|
||||
apply hfg.ne
|
||||
funext j
|
||||
simpa using congrFun h j.succ
|
||||
|
||||
/-- The maximal chain in `Fin n → β`: walk the first tuple element from `⊥` to `⊤`
|
||||
through `β`'s longest chain, then do that with the second element, and so on. -/
|
||||
private def stdChain : (n : ℕ) →
|
||||
{ s : LTSeries (Fin n → β) //
|
||||
s.head = (⊥ : Fin n → β) ∧
|
||||
s.length = n * (FiniteHeightLattice.longestChain (α := β)).length }
|
||||
| 0 => ⟨RelSeries.singleton _ ⊥, by rw [RelSeries.head_singleton], by simp⟩
|
||||
| n + 1 =>
|
||||
let prev := stdChain n
|
||||
⟨RelSeries.smash
|
||||
((FiniteHeightLattice.longestChain (α := β)).map
|
||||
(fun b => (Fin.cons b (⊥ : Fin n → β) : Fin (n + 1) → β)) consBot_strictMono)
|
||||
(prev.1.map (fun f => (Fin.cons (⊤ : β) f : Fin (n + 1) → β)) consTop_strictMono)
|
||||
(by rw [LTSeries.last_map, LTSeries.head_map, prev.2.1]; rfl),
|
||||
by
|
||||
simp only [RelSeries.head_smash, LTSeries.head_map]
|
||||
rw [show (FiniteHeightLattice.longestChain (α := β)).head = (⊥ : β) from rfl]
|
||||
funext i
|
||||
refine Fin.cases ?_ (fun j => ?_) i <;> simp [Pi.bot_apply],
|
||||
by
|
||||
show (FiniteHeightLattice.longestChain (α := β)).length + prev.1.length
|
||||
= (n + 1) * (FiniteHeightLattice.longestChain (α := β)).length
|
||||
rw [prev.2.2, Nat.succ_mul]; exact Nat.add_comm _ _⟩
|
||||
|
||||
instance instFiniteHeight {n : ℕ} : FiniteHeightLattice (Fin n → β) where
|
||||
toLattice := inferInstance
|
||||
longestChain := (stdChain n).1
|
||||
toOrderBot := inferInstance
|
||||
toOrderTop := inferInstance
|
||||
height := n * FiniteHeightLattice.height (α := β)
|
||||
chains_bounded := fun c => by
|
||||
obtain ⟨cs, _, _, hbound⟩ := exists_unzip c
|
||||
refine hbound.trans ?_
|
||||
rw [(stdChain n).2.2]
|
||||
calc ∑ i, (cs i).length
|
||||
≤ ∑ _i : Fin n, (FiniteHeightLattice.longestChain (α := β)).length :=
|
||||
≤ ∑ _i : Fin n, FiniteHeightLattice.height (α := β) :=
|
||||
Finset.sum_le_sum (fun i _ => FiniteHeightLattice.chains_bounded (cs i))
|
||||
_ = n * (FiniteHeightLattice.longestChain (α := β)).length := by
|
||||
_ = n * FiniteHeightLattice.height (α := β) := by
|
||||
simp [Finset.sum_const, Finset.card_univ, Fintype.card_fin]
|
||||
|
||||
end FiniteHeight
|
||||
|
||||
102
lean/Spa/Transformation/Licm.lean
Normal file
102
lean/Spa/Transformation/Licm.lean
Normal file
@@ -0,0 +1,102 @@
|
||||
import Spa.Analysis.Reaching
|
||||
import Spa.Language.Tagged.Graphs
|
||||
|
||||
/-!
|
||||
# Finding loop-invariant assignments (LICM groundwork)
|
||||
|
||||
This wires the **reaching-definitions** analysis (`Spa/Analysis/Reaching.lean`)
|
||||
to the **tagged AST** to *find* — not yet move — assignments inside a `while`
|
||||
loop whose right-hand side depends only on definitions made *outside* the loop.
|
||||
These are the candidates a later LICM pass could hoist.
|
||||
|
||||
The pipeline, for each assignment immediately enclosed by a loop:
|
||||
|
||||
1. locate its CFG state via the tagged-graph bridge (`Program.stateOfNodeId`);
|
||||
2. read the reaching definitions at the assignment's *entry*
|
||||
(`joinForKey s result` — the join over predecessors, i.e. before the
|
||||
assignment itself runs);
|
||||
3. union the definition sets of the RHS variables;
|
||||
4. map each definition site back to its `RawId` (`Program.nodeIdOf`) and check
|
||||
it is **not** inside the loop body (structural `subtreeIds` membership).
|
||||
|
||||
If every reaching definition of every RHS variable lies outside the loop, the
|
||||
assignment is reported as loop-invariant. This is the first-order check ("all
|
||||
reaching definitions outside the loop"); transitive/iterated invariance and the
|
||||
actual hoisting are out of scope here.
|
||||
-/
|
||||
|
||||
namespace Spa
|
||||
|
||||
namespace LicmTransformation
|
||||
|
||||
open Forward
|
||||
|
||||
/-- An assignment found inside a loop, paired with the data needed to test its
|
||||
invariance against that (immediately enclosing) loop. -/
|
||||
structure Candidate (prog : Program) where
|
||||
/-- The enclosing `whileLoop`'s tag (for reporting). -/
|
||||
loopId : prog.NodeId
|
||||
/-- Every node id inside the loop body (the "is-child-of-loop" set). -/
|
||||
bodyIds : List prog.NodeId
|
||||
/-- The assignment `BasicStmt`'s tag — what labels its CFG node. -/
|
||||
assignId : prog.NodeId
|
||||
/-- The variables read by the assignment's RHS. -/
|
||||
rhsVars : List String
|
||||
|
||||
/-- Collect every assignment together with its *immediately enclosing* loop.
|
||||
`enclosing` carries the current loop's tag and body id-set, or `none` outside any
|
||||
loop (in which case assignments are skipped — only in-loop assignments are
|
||||
candidates). -/
|
||||
def collectCandidates (prog : Program) (enc : Option (prog.NodeId × List prog.NodeId)) :
|
||||
Stmt.Tagged prog.NodeId → List (Candidate prog)
|
||||
| .basic _ bs =>
|
||||
match bs, enc with
|
||||
| .assign t _ e, some (loopId, bodyIds) =>
|
||||
[{ loopId := loopId, bodyIds := bodyIds, assignId := t,
|
||||
rhsVars := e.erase.vars.sort (· ≤ ·) }]
|
||||
| _, _ => []
|
||||
| .andThen _ a b => collectCandidates prog enc a ++ collectCandidates prog enc b
|
||||
| .ifElse _ _ a b => collectCandidates prog enc a ++ collectCandidates prog enc b
|
||||
| .whileLoop loopT _ body =>
|
||||
collectCandidates prog (some (loopT, body.subtreeIds)) body
|
||||
|
||||
/-- Read the definition set assigned to variable `k`, or `⊥` if absent. -/
|
||||
def lookupDef (prog : Program) (vs : VariableValues (DefSet prog) prog)
|
||||
(k : String) : DefSet prog :=
|
||||
if h : FiniteMap.MemKey k vs then (FiniteMap.locate h).1 else ⊥
|
||||
|
||||
/-- The AST node ids marked as definition sites in a `DefSet` (those mapped to
|
||||
`true`). With the AST-id-keyed lattice these are recovered directly. -/
|
||||
def defSites (prog : Program) (d : DefSet prog) : List prog.NodeId :=
|
||||
(List.finRange prog.size).filter (fun i => d i)
|
||||
|
||||
/-- Is the candidate assignment loop-invariant: do all reaching definitions of
|
||||
its RHS variables lie outside the loop body? Reaching sets are now keyed by AST
|
||||
node id, so we compare against the loop-body ids directly (embedding the raw
|
||||
body ids into `p.NodeId`). -/
|
||||
def isInvariant (prog : Program) (c : Candidate prog) : Bool :=
|
||||
match prog.stateOfNodeId c.assignId with
|
||||
| none => false
|
||||
| some s =>
|
||||
let entry := joinForKey s (result (DefSet prog) prog)
|
||||
let combined : DefSet prog :=
|
||||
c.rhsVars.foldl (fun acc k => acc ⊔ lookupDef prog entry k) ⊥
|
||||
(defSites prog combined).all (fun nid => ! decide (nid ∈ c.bodyIds))
|
||||
|
||||
/-- The loop-invariant assignments of `prog`, as `(loopId, assignId)` pairs. -/
|
||||
def licmCandidates (prog : Program) : List (prog.NodeId × prog.NodeId) :=
|
||||
(collectCandidates prog none prog.taggedFin).filterMap (fun c =>
|
||||
if isInvariant prog c then some (c.loopId, c.assignId) else none)
|
||||
|
||||
/-- A human-readable report of the loop-invariant assignments. -/
|
||||
def output (prog : Program) : String :=
|
||||
match licmCandidates prog with
|
||||
| [] => "no loop-invariant assignments found"
|
||||
| cands =>
|
||||
"loop-invariant assignments (loop ↦ assignment):\n" ++
|
||||
String.intercalate "\n"
|
||||
(cands.map (fun p => s!" loop #{p.1.val}: assignment #{p.2.val}"))
|
||||
|
||||
end LicmTransformation
|
||||
|
||||
end Spa
|
||||
Reference in New Issue
Block a user