18 Commits

Author SHA1 Message Date
2234f4d0f8 Conclude restricted LICM proof 2026-10-06 21:11:45 -05:00
141fe5dc9d Add intermediate proof that rhs of hoisted assignments are unchanged 2026-10-06 20:44:26 -05:00
13008121d5 Implement guarded LICM transformation 2026-10-06 20:22:35 -05:00
ac99bc047a Show that if x is not assigned within a segment, its value remains as before 2026-10-06 20:16:35 -05:00
cfcd3948a3 Proof step and variable lemmas
1. no writes = save value
2. embedding commutation with steps
3. all variables in code end up in the set of vars
2026-10-06 20:09:08 -05:00
f55a440784 Add environment, expr, and stmt equivalence lemmas 2026-10-06 19:44:57 -05:00
655b7de684 Switch steps to not redundantly include code 2026-10-06 19:30:35 -05:00
234d17394e Prove correctness of constant folding 2026-10-04 10:27:37 -05:00
d2b6bf5af7 Use a unified representation for all trace types 2026-10-04 09:58:39 -05:00
7c05adadff Restore constant folding and expression correctness without trace experiments 2026-09-29 21:04:47 -05:00
53f8bd47dc Add function back in to Embedding 2026-08-09 21:43:20 -05:00
fd371ba175 Remove trace suffix from Reaching type 2026-08-09 21:23:21 -05:00
df4d072f22 Clean up comments in Graphs.lean and Program.lean 2026-08-09 18:17:02 -05:00
c0542d0811 Allow negative numbers in expressions 2026-08-09 17:51:58 -05:00
a19f9fa148 Get rid of Tagged 2026-08-09 17:38:56 -05:00
269906871f Update LICM/Reaching to node use NodeId 2026-08-09 17:35:30 -05:00
1eecf45c0f Add more machinery to use embeddings as "proofs of child-ship" 2026-08-09 17:30:38 -05:00
827d55c6b6 Switch embeddings to index-offset.
This is a special case of an embedding, but it has the nice
property for checking inclusion.
2026-08-09 17:23:46 -05:00
22 changed files with 1321 additions and 1134 deletions

View File

@@ -7,9 +7,11 @@ import Spa.Lattice.Bool
import Spa.Language.Base import Spa.Language.Base
import Spa.Language.Notation import Spa.Language.Notation
import Spa.Language.Semantics import Spa.Language.Semantics
import Spa.Language.Equivalence
import Spa.Language.Graphs import Spa.Language.Graphs
import Spa.Language.Traces import Spa.Language.Traces
import Spa.Language.Properties import Spa.Language.Properties
import Spa.Language.TraceProperties
import Spa.Language import Spa.Language
import Spa.Analysis.Forward.Lattices import Spa.Analysis.Forward.Lattices
import Spa.Analysis.Forward.Evaluation import Spa.Analysis.Forward.Evaluation
@@ -19,10 +21,8 @@ import Spa.Showable
import Spa.Analysis.Utils import Spa.Analysis.Utils
import Spa.Analysis.Sign import Spa.Analysis.Sign
import Spa.Analysis.Constant 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.Analysis.Reaching
import Spa.Analysis.Reaching.Paths
import Spa.Transformation.Licm import Spa.Transformation.Licm
import Spa.Transformation.Licm.Correctness
import Spa.Transformation.Constant

View File

@@ -137,12 +137,11 @@ 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 Forward.analyze_correct ConstLattice prog hrun
theorem analyze_correct_at {ρf : Env} (hrun : EvalStmt [] prog.rootStmt ρf) theorem analyze_correct_at {s : prog.State} {ρin ρout : Env}
{s : prog.State} {ρin ρout : Env} (hr : Reaches s ρin ρout) :
(hr : Reaches (prog.trace hrun) s ρin ρout) :
⟦ joinForKey s (result ConstLattice prog) ⟧ ρin ⟦ joinForKey s (result ConstLattice prog) ⟧ ρin
∧ ⟦ variablesAt s (result ConstLattice prog) ⟧ ρout := ∧ ⟦ variablesAt s (result ConstLattice prog) ⟧ ρout :=
Forward.analyze_correct_at ConstLattice prog hrun hr Forward.analyze_correct_at ConstLattice prog hr
end ConstAnalysis end ConstAnalysis

View File

@@ -90,75 +90,57 @@ lemma stepTrace {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env}
rw [variablesAt_joinAll] rw [variablesAt_joinAll]
exact hjoin exact hjoin
/-- Soundness at *every* visited node: if the analysis result over-approximates the /-- Soundness propagates along an execution prefix: if the analysis is sound at
incoming environment at the start of the trace, then at each node reached along the `s₂` for the run so far (`trₗ`), then it is sound wherever the further prefix
way it over-approximates both the environment entering that node (via `joinForKey`) `mid` ends up. -/
and the environment leaving it (via `variablesAt`). The intermediate `variablesAt` lemma walkPrefix : ∀ {s₂ s : prog.State} {ρ₂ ρin : Env}
evidence used to be computed and discarded inside `walkTrace`; here it is returned. -/ (mid : Traceₗ prog.cfg s₂ s ρ₂ ρin) {s₁ : prog.State} {ρ₁ : Env}
lemma walkTrace_reaches {s₁ s₂ s₃: prog.State} {ρ₁ ρ₂ ρ₃: Env} (trₗ : Traceₗ prog.cfg s₁ s₂ ρ₁ ρ₂),
{s : prog.State} {ρin ρout : Env} ⟦ joinForKey s₂ (result L prog) ⟧ (S.Pre trₗ) →
{tr : Trace prog.cfg s₂ s₃ ρ₂ ρ₃} ⟦ joinForKey s (result L prog) ⟧ (S.Pre (trₗ ++ mid)) := by
(hr : Reaches tr s ρin ρout) intro s₂ s ρ₂ ρin mid
(trₗ : Traceₗ prog.cfg s₁ s₂ ρ₁ ρ₂) match mid with
(hjoin : ⟦ joinForKey s₂ (result L prog) ⟧ (S.Pre trₗ)) : | Traceₗ.nil =>
⟦ joinForKey s (result L prog) ⟧ (S.Pre (trₗ ++ hr.pre)) intro s₁ ρ₁ trₗ hjoin
∧ ⟦ variablesAt s (result L prog) ⟧ (S.Post (trₗ ++ hr.post)) := by simpa only [HAppend.hAppend, Path.append_nil] using hjoin
induction hr with | Traceₗ.cons hnode hedge rest =>
| single_here hnode => intro s₁ ρ₁ trₗ hjoin
simp [Reaches.pre, Reaches.post]
refine ⟨?_, ?_⟩ <;> try simpa [HAppend.hAppend]
exact stepTrace trₗ hjoin hnode
| edge_here hnode hedge rest =>
simp [Reaches.pre, Reaches.post]
refine ⟨?_, ?_⟩ <;> try simpa [HAppend.hAppend]
exact stepTrace trₗ hjoin hnode
| edge_there hnode hedge rest hr' ih =>
have hstep := stepTrace trₗ hjoin hnode have hstep := stepTrace trₗ hjoin hnode
have hmem := FiniteMap.mem_valuesAt prog.states_nodup have hmem := FiniteMap.mem_valuesAt prog.states_nodup
(prog.mem_incoming_of_edge hedge) (variablesAt_mem _ (result L prog)) (prog.mem_incoming_of_edge hedge) (variablesAt_mem _ (result L prog))
simpa [Reaches.pre, Reaches.post, HAppend.hAppend] using simpa only [HAppend.hAppend, Traceₗ.appendStep, Trace.addEdge,
ih ((trₗ ++ hnode).addEdge hedge) Path.append_assoc, Path.single, Path.append] using
walkPrefix rest ((trₗ ++ hnode).addEdge hedge)
(interp_foldr (S.post_pre (trₗ ++ hnode) hedge hstep) hmem) (interp_foldr (S.post_pre (trₗ ++ hnode) hedge hstep) hmem)
omit [DecidableEq L] in omit [DecidableEq L] in
/-- The final node of a trace is always reached, with the environment/state the trace /-- The final node of a trace is always reached, with the environment/state the trace
ends in. Used to recover the final-state soundness theorem from `walkTrace_reaches`. -/ ends in. Used to recover the final-state soundness theorem from `walkPrefix`. -/
def reaches_final {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env} def reaches_final {s : prog.State} {ρ : Env}
(tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂) : (tr : Trace prog.cfg prog.initialState s [] ρ) : Σ ρin, Reaches s ρin ρ :=
Σ ρin, Reaches tr s₂ ρin ρ₂ := ⟨_, ⟨tr.split.2.1, tr.split.2.2⟩⟩
match tr with
| .single hnode => ⟨_, .single_here hnode⟩
| .edge hnode hedge rest =>
let ⟨ρin, r'⟩ := reaches_final rest; ⟨ρin, .edge_there hnode hedge _ r'⟩
omit [DecidableEq L] in omit [DecidableEq L] in
/-- Reaching the final node covers the whole trace. -/ @[simp] lemma reaches_final_post {s : prog.State} {ρ : Env}
@[simp] lemma reaches_final_post {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env} (tr : Trace prog.cfg prog.initialState s [] ρ) :
(tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂) : (reaches_final tr).2.post = tr := Trace.split_append tr
(reaches_final tr).2.post = tr := by
induction tr with
| single hnode => rfl
| edge hnode hedge rest ih => simp [reaches_final, Reaches.post, ih]
variable (L prog) in variable (L prog) in
/-- Soundness at every program point reached during execution: for any node `s` visited /-- Soundness at every program point an execution actually visits: the analysis
by the run `hrun` (witnessed by `hr`), the analysis result over-approximates both the over-approximates both the environment entering that point and the one leaving
environment entering `s` and the one leaving it. The final-state theorem it. -/
`analyze_correct_state` is the special case where `s` is `prog.finalState`. -/ theorem analyze_correct_at {s : prog.State} {ρin ρout : Env} (hr : Reaches s ρin ρout) :
theorem analyze_correct_at {ρf : Env} (hrun : EvalStmt [] prog.rootStmt ρf)
{s : prog.State} {ρin ρout : Env}
(hr : Reaches (prog.trace hrun) s ρin ρout) :
⟦ joinForKey s (result L prog) ⟧ (S.Pre hr.pre) ⟦ joinForKey s (result L prog) ⟧ (S.Pre hr.pre)
∧ ⟦ variablesAt s (result L prog) ⟧ (S.Post hr.post) := by ∧ ⟦ variablesAt s (result L prog) ⟧ (S.Post hr.post) :=
refine walkTrace_reaches hr (Traceₗ.single _ _ []) ?_ have hpre := walkPrefix hr.pre Traceₗ.nil
rw [joinForKey_initialState] (by rw [joinForKey_initialState]; exact ValidStateEvaluator.botV_init)
exact ValidStateEvaluator.botV_init ⟨hpre, stepTrace hr.pre hpre hr.step⟩
variable (L prog) in variable (L prog) in
theorem analyze_correct' theorem analyze_correct'
{ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) : {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
⟦ variablesAt prog.finalState (result L prog) ⟧ (S.Post (prog.trace hrun)) := by ⟦ variablesAt prog.finalState (result L prog) ⟧ (S.Post (prog.trace hrun)) := by
have h := (analyze_correct_at L prog hrun (reaches_final (prog.trace hrun)).2).2 have h := (analyze_correct_at L prog (reaches_final (prog.trace hrun)).2).2
rwa [reaches_final_post] at h rwa [reaches_final_post] at h
end end

View File

@@ -1,6 +1,5 @@
import Spa.Analysis.Forward import Spa.Analysis.Forward
import Spa.Lattice.Finset import Spa.Lattice.Finset
import Spa.Language.Tagged.Graphs
import Spa.Showable import Spa.Showable
namespace Spa namespace Spa
@@ -13,20 +12,18 @@ instance {n : ℕ} : Showable (Finset (Fin n)) :=
(fun i rest => if i ∈ s then show' i ++ ", " ++ rest else rest) "" (fun i rest => if i ∈ s then show' i ++ ", " ++ rest else rest) ""
++ "}"⟩ ++ "}"⟩
abbrev DefSet (prog : Program) : Type := Finset prog.NodeId abbrev DefSet (prog : Program) : Type := Finset prog.State
namespace ReachingAnalysis namespace ReachingAnalysis
variable (prog : Program) variable (prog : Program)
def genSet (s : prog.State) : DefSet prog := (prog.nodeIdOf s).elim {} (fun x => {x})
def eval (s : prog.State) (vs : VariableValues (DefSet prog) prog) : VariableValues (DefSet prog) prog := def eval (s : prog.State) (vs : VariableValues (DefSet prog) prog) : VariableValues (DefSet prog) prog :=
match prog.code s with match prog.code s with
| none => vs | none => vs
| some bs => | some bs =>
match bs with match bs with
| .assign k _ => FiniteMap.generalizedUpdate id (fun _ _ => genSet prog s) [k] vs | .assign k _ => FiniteMap.generalizedUpdate id (fun _ _ => {s}) [k] vs
| .noop => vs | .noop => vs
lemma eval_mono (s : prog.State) : lemma eval_mono (s : prog.State) :
@@ -43,37 +40,37 @@ instance stmtEvaluator : StmtEvaluator (DefSet prog) prog :=
def output : String := def output : String :=
show' (result (DefSet prog) prog) show' (result (DefSet prog) prog)
/-- The statements a trace executed, paired with the state each executed at, /-- Executed nodes, most recent first. Instructions are read from `prog.code`.
most recent first (matching `LastAssign`, which scans for the most recent This is `Path.steps` (chronological) reversed, so facts about concatenating
assignment). This is `Trace.steps` (chronological) reversed, so facts about traces reduce to mathlib's `List.append`/`List.reverse` lemmas. -/
concatenating traces reduce to mathlib's `List.append`/`List.reverse` lemmas. -/ abbrev Run (prog : Program) : Type := List prog.State
abbrev Run (prog : Program) : Type := List (prog.State × BasicStmt)
/-- The first node in a newest-first history whose instruction assigns `x`. -/
@[aesop unsafe cases] @[aesop unsafe cases]
inductive LastAssign (prog : Program) (x : String) : Run prog → prog.NodeId → Prop inductive LastAssign (prog : Program) (x : String) : Run prog → prog.State → Prop
| here (s : prog.State) (e : Expr) (hc : prog.code s = some (.assign x e)) | here (s : prog.State) (e : Expr) (rest : Run prog)
(rest : Run prog) : (hc : prog.code s = some (.assign x e)) :
LastAssign prog x ((s, .assign x e) :: rest) (prog.nodeIdOfNonempty s hc) LastAssign prog x (s :: rest) s
| there (s : prog.State) (bs : BasicStmt) (hc : prog.code s = some bs) | there (s : prog.State) (rest : Run prog) {n : prog.State} :
(rest : Run prog) {n : prog.NodeId} : (∀ e, prog.code s ≠ some (.assign x e)) → LastAssign prog x rest n →
(∀ e, bs ≠ .assign x e) → LastAssign prog x rest n → LastAssign prog x (s :: rest) n
LastAssign prog x ((s, bs) :: rest) n
def runOfTraceₗ {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env} def runOfPath {a b : Configuration prog.cfg} (p : Path prog.cfg a b) : Run prog :=
(tr : Traceₗ prog.cfg s₁ s₂ ρ₁ ρ₂) : Run prog := p.steps.reverse
tr.steps.reverse
def runOfTrace {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env} abbrev runOfTraceₗ {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env}
(tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂) : Run prog := (tr : Traceₗ prog.cfg s₁ s₂ ρ₁ ρ₂) : Run prog := runOfPath prog tr
tr.steps.reverse
abbrev runOfTrace {s₁ s₂ : prog.State} {ρ₁ ρ₂ : Env}
(tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂) : Run prog := runOfPath prog tr
instance stateInterp : StateInterpretation (DefSet prog) prog where instance stateInterp : StateInterpretation (DefSet prog) prog where
Proj := Run prog Proj := Run prog
Pre := @runOfTraceₗ prog Pre := fun tr => runOfPath prog tr
Post := @runOfTrace prog Post := fun tr => runOfPath prog tr
interp vs run := ∀ (x : String) (assigners : DefSet prog), (x, assigners) ∈ vs → interp vs run := ∀ (x : String) (assigners : DefSet prog), (x, assigners) ∈ vs →
∀ (n : prog.NodeId), LastAssign prog x run n → n ∈ assigners ∀ (n : prog.State), LastAssign prog x run n → n ∈ assigners
interp_sup := by interp_sup := by
intro vs₁ vs₂ run h x assigners hmem n hla intro vs₁ vs₂ run h x assigners hmem n hla
obtain ⟨a₁, a₂, rfl, h₁, h₂⟩ := FiniteMap.mem_sup hmem obtain ⟨a₁, a₂, rfl, h₁, h₂⟩ := FiniteMap.mem_sup hmem
@@ -85,37 +82,37 @@ instance stateInterp : StateInterpretation (DefSet prog) prog where
post_pre := by post_pre := by
intro vs s₁ s₂ s₃ ρ₁ ρ₂ tr hedge hvs intro vs s₁ s₂ s₃ ρ₁ ρ₂ tr hedge hvs
simpa [runOfTrace, runOfTraceₗ] using hvs simpa only [runOfPath, Trace.addEdge, Path.steps_append, Path.single,
Path.steps, Step.steps, List.append_nil] using hvs
private lemma valid_step (s : prog.State) {ρ₁ ρ₂ : Env} private lemma valid_step (s : prog.State)
{obs : Option BasicStmt} (hcode : prog.code s = obs)
(hbs : EvalBasicStmtOpt ρ₁ obs ρ₂)
{vs : VariableValues (DefSet prog) prog} {run : Run prog} {vs : VariableValues (DefSet prog) prog} {run : Run prog}
(hvs : ⟦vs⟧ run) : (hvs : ⟦vs⟧ run) :
⟦eval prog s vs⟧ ((hbs.steps s).reverse ++ run) := by ⟦eval prog s vs⟧ ((match prog.code s with | none => [] | some _ => [s]) ++ run) := by
cases hbs with cases hcode : prog.code s with
| none => simpa [eval, hcode, EvalBasicStmtOpt.steps] using hvs | none => simpa [eval, hcode] using hvs
| some hbs => | some bs =>
cases hbs with cases bs with
| noop => | noop =>
simp [eval, hcode, EvalBasicStmtOpt.steps] simp [eval, hcode]
intro x assigners hmem n hla; aesop intro x assigners hmem n hla; aesop (add simp hcode)
| assign x e v hev => | assign x e =>
simp [eval, hcode, EvalBasicStmtOpt.steps]; intro k assigners hmem n hla simp [eval, hcode]; intro k assigners hmem n hla
by_cases hx : k = x by_cases hx : k = x
· subst hx · subst hx
have hd := FiniteMap.generalizedUpdate_mem_eq (List.mem_singleton.mpr rfl) hmem have hd := FiniteMap.generalizedUpdate_mem_eq (List.mem_singleton.mpr rfl) hmem
rcases hla rcases hla <;> simp [hd] <;> aesop (add simp hcode)
<;> simp [Program.nodeIdOfNonempty, hd, genSet, Option.get] <;> aesop
· have hmem' := FiniteMap.generalizedUpdate_not_mem_backward · have hmem' := FiniteMap.generalizedUpdate_not_mem_backward
(fun hc => hx (List.mem_singleton.mp hc)) hmem (fun hc => hx (List.mem_singleton.mp hc)) hmem
aesop aesop (add simp hcode)
instance validStateEvaluator : ValidStateEvaluator (DefSet prog) prog where instance validStateEvaluator : ValidStateEvaluator (DefSet prog) prog where
valid := by valid := by
intro s₁ s₂ ρ₁ ρ₂ ρ₃ vs tr hbs hvs intro s₁ s₂ ρ₁ ρ₂ ρ₃ vs tr hbs hvs
show ⟦eval prog s₂ vs⟧ (runOfTrace prog (tr ++ hbs)) change ⟦vs⟧ (runOfPath prog tr) at hvs
simpa [runOfTrace, runOfTraceₗ] using valid_step prog s₂ rfl hbs hvs change ⟦eval prog s₂ vs⟧ (runOfPath prog (Path.append tr (.single (.execute hbs))))
cases hcode : prog.code s₂ <;>
simpa [runOfPath, Path.single, Path.steps, Step.steps, hcode] using valid_step prog s₂ hvs
botV_init := by intro x assigners _ n hla; cases hla botV_init := by intro x assigners _ n hla; cases hla
theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) : theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
@@ -123,12 +120,11 @@ theorem analyze_correct {ρ : Env} (hrun : EvalStmt [] prog.rootStmt ρ) :
(runOfTrace prog (prog.trace hrun)) := (runOfTrace prog (prog.trace hrun)) :=
Forward.analyze_correct' (DefSet prog) prog hrun Forward.analyze_correct' (DefSet prog) prog hrun
theorem analyze_correct_at {ρf : Env} (hrun : EvalStmt [] prog.rootStmt ρf) theorem analyze_correct_at {s : prog.State} {ρin ρout : Env}
{s : prog.State} {ρin ρout : Env} (hr : Reaches s ρin ρout) :
(hr : Reaches (prog.trace hrun) s ρin ρout) :
⟦ joinForKey s (result (DefSet prog) prog) ⟧ (runOfTraceₗ prog hr.pre) ⟦ joinForKey s (result (DefSet prog) prog) ⟧ (runOfTraceₗ prog hr.pre)
∧ ⟦ variablesAt s (result (DefSet prog) prog) ⟧ (runOfTrace prog hr.post) := ∧ ⟦ variablesAt s (result (DefSet prog) prog) ⟧ (runOfTrace prog hr.post) :=
Forward.analyze_correct_at (DefSet prog) prog hrun hr Forward.analyze_correct_at (DefSet prog) prog hr
end ReachingAnalysis end ReachingAnalysis

View File

@@ -0,0 +1,54 @@
import Spa.Analysis.Reaching
import Spa.Language.TraceProperties
namespace Spa
namespace ReachingAnalysis
/-- The most recent assignment occurs in the history being searched. -/
lemma LastAssign.mem {prog : Program} {x : String} {run : Run prog} {d : prog.State}
(h : LastAssign prog x run d) : d ∈ run := by
induction h <;> aesop
/-- Appending older history cannot displace an already-found assignment. -/
lemma LastAssign.append {prog : Program} {x : String} {new : Run prog} {d : prog.State}
(h : LastAssign prog x new d) (old : Run prog) :
LastAssign prog x (new ++ old) d := by
induction h with
| here s rhs rest hc => exact .here s rhs _ hc
| there s rest hn h ih => exact .there s _ hn ih
/-- A history containing a write to `x` has a most recent assignment to `x`. -/
lemma lastAssign_of_write {prog : Program} {x : String} {run : Run prog}
(hw : ∃ d ∈ run, ∃ rhs, prog.code d = some (.assign x rhs)) :
∃ d, LastAssign prog x run d := by
induction run with
| nil => simp at hw
| cons d rest ih =>
by_cases hx : ∃ rhs, prog.code d = some (.assign x rhs)
· obtain ⟨rhs, hc⟩ := hx
exact ⟨d, .here d rhs rest hc⟩
· have hw' : ∃ j ∈ rest, ∃ rhs, prog.code j = some (.assign x rhs) := by
obtain ⟨j, hm, rhs, hc⟩ := hw
rcases List.mem_cons.mp hm with rfl | hm
· exact False.elim (hx ⟨rhs, hc⟩)
· exact ⟨j, hm, rhs, hc⟩
obtain ⟨j, hj⟩ := ih hw'
exact ⟨j, .there d rest (by simpa using hx) hj⟩
/-- Outside reaching definitions rule out any write in a confined intervening
path. No equality of static sites is used to infer equality of events. -/
lemma Path.preserves_of_lastAssign_outside {prog : Program}
{a b c : Configuration prog.cfg} (pre : Path prog.cfg a b) (seg : Path prog.cfg b c)
{x : String} (sites : Set prog.State)
(hin : ∀ d ∈ seg.steps, d ∈ sites)
(hout : ∀ d, LastAssign prog x (runOfPath prog (pre.append seg)) d → d ∉ sites) :
∀ v, Env.Mem (x, v) b.2 ↔ Env.Mem (x, v) c.2 := by
apply seg.preserves_unwritten
intro d hm rhs hc
obtain ⟨j, hl⟩ := lastAssign_of_write ⟨d, List.mem_reverse.mpr hm, rhs, hc⟩
apply hout j
· simpa [runOfPath, List.reverse_append] using hl.append (runOfPath prog pre)
· exact hin j (List.mem_reverse.mp hl.mem)
end ReachingAnalysis
end Spa

View File

@@ -111,13 +111,16 @@ namespace SignAnalysis
variable (prog : Program) variable (prog : Program)
/-- The sign of an integer literal. -/
def signOf (z : ℤ) : SignLattice :=
if z = 0 then .mk .zero else if 0 < z then .mk .plus else .mk .minus
def eval : Expr → VariableValues SignLattice prog → SignLattice def eval : Expr → VariableValues SignLattice prog → SignLattice
| .add e₁ e₂, vs => plus (eval e₁ vs) (eval e₂ vs) | .add e₁ e₂, vs => plus (eval e₁ vs) (eval e₂ vs)
| .sub e₁ e₂, vs => minus (eval e₁ vs) (eval e₂ vs) | .sub e₁ e₂, vs => minus (eval e₁ vs) (eval e₂ vs)
| .var k, vs => | .var k, vs =>
if h : FiniteMap.MemKey k vs then (FiniteMap.locate h).1 else .top if h : FiniteMap.MemKey k vs then (FiniteMap.locate h).1 else .top
| .num 0, _ => .mk .zero | .num z, _ => signOf z
| .num (_ + 1), _ => .mk .plus
lemma eval_mono (e : Expr) : Monotone (eval prog e) := by lemma eval_mono (e : Expr) : Monotone (eval prog e) := by
induction e with induction e with
@@ -139,7 +142,7 @@ lemma eval_mono (e : Expr) : Monotone (eval prog e) := by
dif_neg (fun hm => hk (FiniteMap.MemKey_iff.mp hm))] dif_neg (fun hm => hk (FiniteMap.MemKey_iff.mp hm))]
| num n => | num n =>
intro vs₁ vs₂ _ intro vs₁ vs₂ _
cases n <;> exact le_refl _ exact le_refl _
instance exprEvaluator : ExprEvaluator SignLattice prog := instance exprEvaluator : ExprEvaluator SignLattice prog :=
⟨eval prog, eval_mono prog⟩ ⟨eval prog, eval_mono prog⟩
@@ -159,6 +162,20 @@ private lemma int_neg_iff (z : ℤ) : (∃ n : ℕ, z = -((n : ℤ) + 1)) ↔ z
· rintro ⟨n, rfl⟩; omega · rintro ⟨n, rfl⟩; omega
· intro h; exact ⟨(-z - 1).toNat, by omega⟩ · intro h; exact ⟨(-z - 1).toNat, by omega⟩
/-- `signOf` really does describe the literal it was computed from. -/
lemma interp_signOf (z : ℤ) : ⟦signOf z⟧ (Value.int z) := by
unfold signOf
split
· case isTrue h => subst h; rfl
· rename_i hne
split
· case isTrue hpos =>
simp only [signInterpretation, interpSign, Value.int.injEq, int_pos_iff]
exact hpos
· case isFalse hnpos =>
simp only [signInterpretation, interpSign, Value.int.injEq, int_neg_iff]
omega
lemma plus_valid {g₁ g₂ : SignLattice} {z₁ z₂ : ℤ} lemma plus_valid {g₁ g₂ : SignLattice} {z₁ z₂ : ℤ}
(h₁ : ⟦g₁⟧ (.int z₁)) (h₂ : ⟦g₂⟧ (.int z₂)) : (h₁ : ⟦g₁⟧ (.int z₁)) (h₂ : ⟦g₂⟧ (.int z₂)) :
⟦plus g₁ g₂⟧ (.int (z₁ + z₂)) := by ⟦plus g₁ g₂⟧ (.int (z₁ + z₂)) := by
@@ -184,9 +201,7 @@ instance eval_valid : ValidExprEvaluator SignLattice prog := by
| num n => | num n =>
intro _ intro _
show ⟦eval prog (.num n) vs⟧ (.int n) show ⟦eval prog (.num n) vs⟧ (.int n)
cases n with exact interp_signOf n
| zero => rfl
| succ n' => exact ⟨n', congrArg Value.int (by norm_cast)⟩
| var x v hxv => | var x v hxv =>
intro hvs intro hvs
show ⟦eval prog (.var x) vs⟧ v show ⟦eval prog (.var x) vs⟧ v
@@ -213,12 +228,11 @@ 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 Forward.analyze_correct SignLattice prog hrun
theorem analyze_correct_at {ρf : Env} (hrun : EvalStmt [] prog.rootStmt ρf) theorem analyze_correct_at {s : prog.State} {ρin ρout : Env}
{s : prog.State} {ρin ρout : Env} (hr : Reaches s ρin ρout) :
(hr : Reaches (prog.trace hrun) s ρin ρout) :
⟦ joinForKey s (result SignLattice prog) ⟧ ρin ⟦ joinForKey s (result SignLattice prog) ⟧ ρin
∧ ⟦ variablesAt s (result SignLattice prog) ⟧ ρout := ∧ ⟦ variablesAt s (result SignLattice prog) ⟧ ρout :=
Forward.analyze_correct_at SignLattice prog hrun hr Forward.analyze_correct_at SignLattice prog hr
end SignAnalysis end SignAnalysis

View File

@@ -5,9 +5,13 @@ import Mathlib.Data.Finset.Basic
# Base Language # Base Language
This file defines the core object language for the program analysis and This file defines the core object language for the program analysis and
transformation. It's a very basic imperative language. The `Spa/Language/Tagged/Basic.lean` transformation. It's a very basic imperative language.
file provides an auto-derived version of the `Expr`, `BasicStmt`, and `Stmt` data
types with unique IDs per condtructor, enabling in-AST pointers. Program points are identified by their node in the control flow graph rather than
by an identifier stored in the AST: a recursion over a `Stmt` threads a
`Spa.GGraph.Embed` of the subtree's CFG into the whole program's (starting from
`Spa.Program.rootEmbed`), which yields the CFG index of each basic statement
along with a proof that the node carries it.
-/ -/
@@ -18,7 +22,7 @@ inductive Expr where
| add (e₁ e₂ : Expr) | add (e₁ e₂ : Expr)
| sub (e₁ e₂ : Expr) | sub (e₁ e₂ : Expr)
| var (x : String) | var (x : String)
| num (n : ℕ) | num (z : ℤ)
deriving DecidableEq deriving DecidableEq
/-- A statement that cannot alter control flow (and thus, can be part of a basic block). /-- A statement that cannot alter control flow (and thus, can be part of a basic block).

View File

@@ -0,0 +1,135 @@
import Spa.Language.Semantics
namespace Spa
/-- Environments agree on the current values of the selected variables. -/
def Env.AgreeOn (xs : Finset String) (ρ σ : Env) : Prop :=
∀ x ∈ xs, ∀ v, Env.Mem (x, v) ρ ↔ Env.Mem (x, v) σ
/-- Observable environment equality, ignoring shadowed bindings. -/
def Env.Equiv (ρ σ : Env) : Prop :=
∀ x v, Env.Mem (x, v) ρ ↔ Env.Mem (x, v) σ
lemma Env.Mem.functional {ρ : Env} {x : String} {v w : Value}
(h : Env.Mem (x, v) ρ) (h' : Env.Mem (x, w) ρ) : v = w := by
induction ρ with
| nil => cases h
| cons pair rest ih =>
cases h <;> cases h' <;> aesop
lemma Env.mem_cons {ρ : Env} {x y : String} {v w : Value} :
Env.Mem (x, v) ((y, w) :: ρ) ↔ (x = y ∧ v = w) ∨ (x ≠ y ∧ Env.Mem (x, v) ρ) := by
constructor
· intro h; cases h <;> aesop
· rintro (⟨rfl, rfl⟩ | ⟨hne, h⟩)
· exact .here _ _ _
· exact .there _ _ _ _ _ hne h
lemma Env.Equiv.refl (ρ : Env) : Env.Equiv ρ ρ := fun _ _ => Iff.rfl
lemma Env.Equiv.symm {ρ σ : Env} (h : Env.Equiv ρ σ) : Env.Equiv σ ρ :=
fun x v => (h x v).symm
lemma Env.Equiv.trans {ρ σ τ : Env} (h : Env.Equiv ρ σ) (h' : Env.Equiv σ τ) :
Env.Equiv ρ τ := fun x v => (h x v).trans (h' x v)
lemma Env.Equiv.cons {ρ σ : Env} (h : Env.Equiv ρ σ) (x : String) (v : Value) :
Env.Equiv ((x, v) :: ρ) ((x, v) :: σ) := by
intro y w; simp only [Env.mem_cons, h y w]
lemma Env.cons_equiv_of_mem {ρ : Env} {x : String} {v : Value}
(h : Env.Mem (x, v) ρ) : Env.Equiv ((x, v) :: ρ) ρ := by
intro y w
rw [Env.mem_cons]
constructor
· rintro (⟨rfl, rfl⟩ | ⟨_, hw⟩) <;> assumption
· intro hw
by_cases he : y = x
· subst y; exact Or.inl ⟨rfl, hw.functional h⟩
· exact Or.inr ⟨he, hw⟩
lemma EvalExpr.congr_env {ρ σ : Env} {e : Expr} {v : Value}
(h : EvalExpr ρ e v) (ha : Env.AgreeOn e.vars ρ σ) : EvalExpr σ e v := by
induction h with
| num => exact .num _ _
| var x v hm => exact .var _ _ _ ((ha x (by simp [Expr.vars]) v).mp hm)
| add a b u v _ _ iha ihb =>
exact .add _ _ _ _ _ (iha (fun x hx => ha x (Finset.mem_union_left _ hx)))
(ihb (fun x hx => ha x (Finset.mem_union_right _ hx)))
| sub a b u v _ _ iha ihb =>
exact .sub _ _ _ _ _ (iha (fun x hx => ha x (Finset.mem_union_left _ hx)))
(ihb (fun x hx => ha x (Finset.mem_union_right _ hx)))
lemma EvalExpr.deterministic {ρ : Env} {e : Expr} {v w : Value}
(h : EvalExpr ρ e v) (h' : EvalExpr ρ e w) : v = w := by
induction h generalizing w with
| num => cases h'; rfl
| var _ _ hm => cases h' with | var _ _ hm' => exact hm.functional hm'
| add a b u v h₁ h₂ ih₁ ih₂ =>
cases h' with
| add _ _ u' v' h₁' h₂' =>
have := ih₁ h₁'; have := ih₂ h₂'; aesop
| sub a b u v h₁ h₂ ih₁ ih₂ =>
cases h' with
| sub _ _ u' v' h₁' h₂' =>
have := ih₁ h₁'; have := ih₂ h₂'; aesop
/-- Variables a statement may assign. -/
def Stmt.writes : Stmt → Finset String
| .basic .noop => ∅
| .basic (.assign x _) => {x}
| .andThen a b => a.writes ∪ b.writes
| .ifElse _ a b => a.writes ∪ b.writes
| .whileLoop _ b => b.writes
lemma EvalStmt.preserves_unwritten {ρ σ : Env} {s : Stmt} (h : EvalStmt ρ s σ)
{x : String} (hx : x ∉ s.writes) :
∀ v, Env.Mem (x, v) ρ ↔ Env.Mem (x, v) σ := by
induction h with
| basic _ _ _ hb =>
cases hb with
| noop => exact fun _ => Iff.rfl
| assign y _ _ _ =>
have hn : x ≠ y := by simpa [Stmt.writes] using hx
intro v; simp [Env.mem_cons, hn]
| andThen _ _ _ _ _ _ _ ih₁ ih₂ =>
simp only [Stmt.writes, Finset.mem_union, not_or] at hx
exact fun v => (ih₁ hx.1 v).trans (ih₂ hx.2 v)
| ifTrue _ _ _ _ _ _ _ _ _ ih =>
exact ih (fun hm => hx (Finset.mem_union_left _ hm))
| ifFalse _ _ _ _ _ _ _ ih =>
exact ih (fun hm => hx (Finset.mem_union_right _ hm))
| whileTrue _ _ _ _ _ _ _ _ _ _ ih₁ ih₂ =>
exact fun v => (ih₁ hx v).trans (ih₂ hx v)
| whileFalse => exact fun _ => Iff.rfl
/-- Evaluations depend on current bindings, not the list of shadowed bindings. -/
noncomputable def EvalStmt.congr_env {ρ ρ' : Env} {s : Stmt} (h : EvalStmt ρ s ρ') :
∀ {σ}, Env.Equiv ρ σ → Σ σ', {_h : EvalStmt σ s σ' // Env.Equiv ρ' σ'} := by
induction h with
| basic _ _ _ hb =>
intro σ he
cases hb with
| noop => exact ⟨σ, .basic _ _ _ (.noop _), he⟩
| assign x rhs v hv =>
exact ⟨_, .basic _ _ _ (.assign _ _ _ _ (hv.congr_env (fun y _ => he y))), he.cons x v⟩
| andThen _ _ _ _ _ _ _ ih₁ ih₂ =>
intro σ he
obtain ⟨σ₁, h₁, he₁⟩ := ih₁ he
obtain ⟨σ₂, h₂, he₂⟩ := ih₂ he₁
exact ⟨σ₂, .andThen _ _ _ _ _ h₁ h₂, he₂⟩
| ifTrue _ _ _ _ _ _ hc hz _ ih =>
intro σ he
obtain ⟨σ', h', he'⟩ := ih he
exact ⟨σ', .ifTrue _ _ _ _ _ _ (hc.congr_env (fun x _ => he x)) hz h', he'⟩
| ifFalse _ _ _ _ _ hc _ ih =>
intro σ he
obtain ⟨σ', h', he'⟩ := ih he
exact ⟨σ', .ifFalse _ _ _ _ _ (hc.congr_env (fun x _ => he x)) h', he'⟩
| whileTrue _ _ _ _ _ _ hc hz _ _ ih₁ ih₂ =>
intro σ he
obtain ⟨σ₁, h₁, he₁⟩ := ih₁ he
obtain ⟨σ₂, h₂, he₂⟩ := ih₂ he₁
exact ⟨σ₂, .whileTrue _ _ _ _ _ _ (hc.congr_env (fun x _ => he x)) hz h₁ h₂, he₂⟩
| whileFalse _ _ _ hc =>
intro σ he
exact ⟨σ, .whileFalse _ _ _ (hc.congr_env (fun x _ => he x)), he⟩
end Spa

View File

@@ -215,62 +215,103 @@ lemma wrap_outputs (g : GGraph (Option β)) :
/-! ### Embeddings /-! ### Embeddings
Each composition operator includes its operands into the result via an index To be able to reason compositionally about traces through the graphs,
translation that preserves node payloads and edges. `Embed` captures exactly we need to be able to reason about how a trace within a sub-graph maps
those two facts, so anything defined from `nodes` and `edges` (traces, node to the full graph. Fortunately, graphs are built using composition operators,
labels, …) can be transported along an embedding once, instead of once per and these composition operators always include their arguments as embedded
operator. subgraphs in the full result. Moreover, each embedding "just" offsets the
existing node IDs by a given amount.
`Embed` is deliberately a structure rather than a class: for `g ⤳ g`, both the This section formalizes this fact by providing an `Embed` type that
left and the right inclusion inhabit the same type `Embed g (g ⤳ g)`, so represents an offset-based embedding, and showing that such an embedding
instance resolution could silently pick the wrong copy. Embeddings into a exists for all arguments given to graph composition operators. Furthermore,
composed graph are non-canonical by design; a named witness says which because of the offset-based embedding, we can determine whether a node
inclusion is meant. -/ came from a particular subgraph simply by examining its offset and sub-graph
size. This is captured by `Embed.mem_range_iff`. -/
/-- An embedding of graph `g` into graph `h`: an index translation that /-- A special-case embedding of `g` into `h` in which all edges and nodes
preserves node payloads and edges. -/ of `g` are present in `h` at a given offset `off`. -/
structure Embed (g h : GGraph α) where structure Embed (g h : GGraph α) where
f : g.Index → h.Index f : g.Index → h.Index
off : ℕ
f_val : ∀ i, (f i).val = off + i.val
nodes_eq : ∀ i, h.nodes (f i) = g.nodes i nodes_eq : ∀ i, h.nodes (f i) = g.nodes i
edges_mem : ∀ {e : g.Edge}, e ∈ g.edges → (f e.1, f e.2) ∈ h.edges edges_mem : ∀ {e : g.Edge}, e ∈ g.edges → (f e.1, f e.2) ∈ h.edges
/-- Embeddings compose. -/ lemma Embed.f_inj {g h : GGraph α} (e : Embed g h) : Function.Injective e.f := by
intro i j hij
have := congrArg Fin.val hij
rw [e.f_val, e.f_val] at this
exact Fin.ext (by omega)
/-- An embedding's range is the interval `[off, off + g.size)`. -/
lemma Embed.mem_range_iff {g h : GGraph α} (e : Embed g h) (j : h.Index) :
(∃ i, e.f i = j) ↔ e.off ≤ j.val ∧ j.val < e.off + g.size := by
constructor
· rintro ⟨i, rfl⟩; have := i.isLt; rw [e.f_val]; omega
· rintro ⟨hlo, hhi⟩
refine ⟨⟨j.val - e.off, by omega⟩, Fin.ext ?_⟩
rw [e.f_val]
show e.off + (j.val - e.off) = j.val
omega
/-- Build an embedding from an index map that is pointwise the shift. The five
inclusions below are naturally written with `Fin.castAdd`/`Fin.natAdd` — the form
the `Fin.append` lemmas are stated in — so this lets them keep those proofs
verbatim. The trailing argument is boilerplate at every call site and defaults
to discharging itself. -/
private def Embed.ofIndexMap {g h : GGraph α} (off : ℕ) (k : g.Index → h.Index)
(hn : ∀ i, h.nodes (k i) = g.nodes i)
(hem : ∀ {e : g.Edge}, e ∈ g.edges → (k e.1, k e.2) ∈ h.edges)
(hk : ∀ i, (k i).val = off + i.val := by intro i; simp) :
Embed g h where
f := k
off := off
f_val := hk
nodes_eq := hn
edges_mem := hem
/-- Embeddings compose (offsets add). -/
def Embed.trans {g₁ g₂ g₃ : GGraph α} (e₁ : Embed g₁ g₂) (e₂ : Embed g₂ g₃) : def Embed.trans {g₁ g₂ g₃ : GGraph α} (e₁ : Embed g₁ g₂) (e₂ : Embed g₂ g₃) :
Embed g₁ g₃ where Embed g₁ g₃ :=
f := e₂.f ∘ e₁.f ofIndexMap (e₂.off + e₁.off) (fun i => e₂.f (e₁.f i))
nodes_eq i := (e₂.nodes_eq (e₁.f i)).trans (e₁.nodes_eq i) (fun i => (e₂.nodes_eq (e₁.f i)).trans (e₁.nodes_eq i))
edges_mem he := e₂.edges_mem (e₁.edges_mem he) (fun he => e₂.edges_mem (e₁.edges_mem he))
(hk := fun i => by rw [e₂.f_val, e₁.f_val]; omega)
/-- The left operand's inclusion into a sequenced graph. -/ /-- The left operand's inclusion into a sequenced graph. -/
def Embed.sequenceLeft (g₁ g₂ : GGraph α) : Embed g₁ (g₁ ⤳ g₂) where def Embed.sequenceLeft (g₁ g₂ : GGraph α) : Embed g₁ (g₁ ⤳ g₂) :=
f i := i.castAdd g₂.size ofIndexMap 0 (fun i => i.castAdd g₂.size) (Fin.append_left g₁.nodes g₂.nodes)
nodes_eq i := Fin.append_left g₁.nodes g₂.nodes i (fun he => List.mem_append_left _ (List.mem_append_left _ (List.mem_map_of_mem _ he)))
edges_mem he := List.mem_append_left _ (List.mem_append_left _ (List.mem_map_of_mem _ he))
/-- The right operand's inclusion into a sequenced graph. -/ /-- The right operand's inclusion into a sequenced graph. -/
def Embed.sequenceRight (g₁ g₂ : GGraph α) : Embed g₂ (g₁ ⤳ g₂) where def Embed.sequenceRight (g₁ g₂ : GGraph α) : Embed g₂ (g₁ ⤳ g₂) :=
f i := i.natAdd g₁.size ofIndexMap g₁.size (fun i => i.natAdd g₁.size) (Fin.append_right g₁.nodes g₂.nodes)
nodes_eq i := Fin.append_right g₁.nodes g₂.nodes i (fun he => List.mem_append_left _ (List.mem_append_right _ (List.mem_map_of_mem _ he)))
edges_mem he := List.mem_append_left _ (List.mem_append_right _ (List.mem_map_of_mem _ he))
/-- The left operand's inclusion into an overlaid graph. -/ /-- The left operand's inclusion into an overlaid graph. -/
def Embed.overlayLeft (g₁ g₂ : GGraph α) : Embed g₁ (g₁ ∙ g₂) where def Embed.overlayLeft (g₁ g₂ : GGraph α) : Embed g₁ (g₁ ∙ g₂) :=
f i := i.castAdd g₂.size ofIndexMap 0 (fun i => i.castAdd g₂.size) (Fin.append_left g₁.nodes g₂.nodes)
nodes_eq i := Fin.append_left g₁.nodes g₂.nodes i (fun he => List.mem_append_left _ (List.mem_map_of_mem _ he))
edges_mem he := List.mem_append_left _ (List.mem_map_of_mem _ he)
/-- The right operand's inclusion into an overlaid graph. -/ /-- The right operand's inclusion into an overlaid graph. -/
def Embed.overlayRight (g₁ g₂ : GGraph α) : Embed g₂ (g₁ ∙ g₂) where def Embed.overlayRight (g₁ g₂ : GGraph α) : Embed g₂ (g₁ ∙ g₂) :=
f i := i.natAdd g₁.size ofIndexMap g₁.size (fun i => i.natAdd g₁.size) (Fin.append_right g₁.nodes g₂.nodes)
nodes_eq i := Fin.append_right g₁.nodes g₂.nodes i (fun he => List.mem_append_right _ (List.mem_map_of_mem _ he))
edges_mem he := List.mem_append_right _ (List.mem_map_of_mem _ he)
/-- The body's inclusion into a `loop` graph. -/ /-- The body's inclusion into a `loop` graph. -/
def Embed.loop (g : GGraph (Option β)) : Embed g (loop g) where def Embed.loop (g : GGraph (Option β)) : Embed g (GGraph.loop g) :=
f i := i.natAdd 2 ofIndexMap 2 (fun i => i.natAdd 2) (Fin.append_right (fun _ : Fin 2 => none) g.nodes)
nodes_eq i := Fin.append_right (fun _ : Fin 2 => none) g.nodes i (fun he => List.mem_append_left _ (List.mem_append_left _
edges_mem he := List.mem_append_left _ (List.mem_append_left _ (List.mem_append_left _ (List.mem_map_of_mem _ he))))
(List.mem_append_left _ (List.mem_map_of_mem _ he)))
/-- A `singleton` subgraph has exactly one node; this is where it sits in the ambient graph. -/
def Embed.singletonIndex {a : α} {h : GGraph α} (e : Embed (singleton a) h) : h.Index :=
e.f ⟨0, Nat.zero_lt_one⟩
@[simp] lemma Embed.nodes_singletonIndex {a : α} {h : GGraph α}
(e : Embed (singleton a) h) : h.nodes e.singletonIndex = a :=
e.nodes_eq ⟨0, Nat.zero_lt_one⟩
variable (g : GGraph α) variable (g : GGraph α)

View File

@@ -31,6 +31,11 @@ def cfg : Graph := Graph.wrap p.rootStmt.cfg
/-- A state in the control flow `Spa.Graph` of this program. -/ /-- A state in the control flow `Spa.Graph` of this program. -/
abbrev State : Type := p.cfg.Index abbrev State : Type := p.cfg.Index
/-- The root statement's CFG sits inside the program's CFG. -/
def rootEmbed : GGraph.Embed p.rootStmt.cfg p.cfg :=
(GGraph.Embed.sequenceLeft p.rootStmt.cfg (Graph.singleton none)).trans
(GGraph.Embed.sequenceRight (Graph.singleton none) _)
/-- Variables mentioned or defined in this program. -/ /-- Variables mentioned or defined in this program. -/
def vars : List String := p.rootStmt.vars.sort (· ≤ ·) def vars : List String := p.rootStmt.vars.sort (· ≤ ·)

View File

@@ -27,16 +27,6 @@ section Embeddings
variable {g₁ g₂ : Graph} {ρ₁ ρ₂ : Env} variable {g₁ g₂ : Graph} {ρ₁ ρ₂ : Env}
/-- Transport a trace along a graph embedding: an embedding preserves node
payloads and edges, which is everything a trace is made of. This is the
single induction behind all the per-operator lifting corollaries below. -/
noncomputable def Trace.embed {g h : Graph} (e : GGraph.Embed g h)
{idx₁ idx₂ : g.Index} (tr : Trace g idx₁ idx₂ ρ₁ ρ₂) :
Trace h (e.f idx₁) (e.f idx₂) ρ₁ ρ₂ := by
induction tr with
| single hbs => exact Trace.single (by rwa [e.nodes_eq])
| edge hbs he _ ih => exact Trace.edge (by rwa [e.nodes_eq]) (e.edges_mem he) ih
/-- When two graphs are overlaid, for each trace in the left graph, /-- When two graphs are overlaid, for each trace in the left graph,
a corresponding trace exists in the combined graph. -/ a corresponding trace exists in the combined graph. -/
noncomputable def Trace.overlay_left {idx₁ idx₂ : g₁.Index} noncomputable def Trace.overlay_left {idx₁ idx₂ : g₁.Index}
@@ -81,6 +71,15 @@ noncomputable def EndToEndTrace.overlay_right (etr : EndToEndTrace g₂ ρ₁ ρ
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⟩ tr.overlay_right⟩
/-- Execute the left operand and follow the connecting edge to the right operand. -/
noncomputable def EndToEndTrace.beforeRight {ρ₃ : Env}
(left : EndToEndTrace g₁ ρ₁ ρ₂) (right : EndToEndTrace g₂ ρ₂ ρ₃) :
Traceₗ (g₁ ⤳ g₂) (left.entry.castAdd g₂.size) (right.entry.natAdd g₁.size) ρ₁ ρ₂ := by
refine left.trace.sequence_left.addEdge ?_
exact List.mem_append_right _
(List.mem_product.mpr
⟨List.mem_map_of_mem _ left.exit_mem, List.mem_map_of_mem _ right.entry_mem⟩)
/-- When two graphs are sequenced, two end-to-end traces through the respective /-- 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 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 graph. This is only possible for end-to-end traces and not for general
@@ -90,13 +89,10 @@ noncomputable def EndToEndTrace.overlay_right (etr : EndToEndTrace g₂ ρ₁ ρ
with a trace in another graph. -/ with a trace in another graph. -/
noncomputable def EndToEndTrace.concat {ρ₃ : Env} (etr₁ : EndToEndTrace g₁ ρ₁ ρ₂) noncomputable def EndToEndTrace.concat {ρ₃ : Env} (etr₁ : EndToEndTrace g₁ ρ₁ ρ₂)
(etr₂ : EndToEndTrace g₂ ρ₂ ρ₃) : EndToEndTrace (g₁ ⤳ g₂) ρ₁ ρ₃ := by (etr₂ : EndToEndTrace g₂ ρ₂ ρ₃) : EndToEndTrace (g₁ ⤳ g₂) ρ₁ ρ₃ := by
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁ exact ⟨etr₁.entry.castAdd g₂.size, List.mem_map_of_mem _ etr₁.entry_mem,
obtain ⟨j₁, k₁, j₂, k₂, tr₂⟩ := etr₂ etr₂.exit.natAdd g₁.size, List.mem_map_of_mem _ etr₂.exit_mem,
refine ⟨i₁.castAdd g₂.size, List.mem_map_of_mem _ h₁, (etr₁.beforeRight etr₂).appendTrace etr₂.trace.sequence_right⟩
j₂.natAdd g₁.size, List.mem_map_of_mem _ k₂,
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 end Embeddings
@@ -119,21 +115,23 @@ private lemma loop_nodes_at_out :
(Graph.loop g).nodes g.loopOut = none := (Graph.loop g).nodes g.loopOut = none :=
Fin.append_left (fun _ : Fin 2 => none) g.nodes 1 Fin.append_left (fun _ : Fin 2 => none) g.nodes 1
/-- Execute the empty loop header and follow its edge into this body execution. -/
noncomputable def EndToEndTrace.beforeBody (body : EndToEndTrace g ρ₁ ρ₂) :
Traceₗ (Graph.loop g) g.loopIn (body.entry.natAdd 2) ρ₁ ρ₁ := by
refine (Trace.single (loop_nodes_at_in ▸ EvalBasicStmtOpt.none)).addEdge ?_
refine List.mem_append_left _ (List.mem_append_left _ (List.mem_append_right _ ?_))
exact List.mem_map_of_mem _ (List.mem_map_of_mem _ body.entry_mem)
/-- Equivlaent of `Trace.loop` for end-to-end traces. -/ /-- Equivlaent of `Trace.loop` for end-to-end traces. -/
noncomputable def EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) : noncomputable def EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) :
EndToEndTrace (Graph.loop g) ρ₁ ρ₂ := by EndToEndTrace (Graph.loop g) ρ₁ ρ₂ := by
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := etr -- the edge (2 ↑ʳ etr.exit) → out, reached through the third edge group
-- the edge in → (2 ↑ʳ i₁), reached through the second edge group have hout : (etr.exit.natAdd 2, g.loopOut) ∈ (Graph.loop g).edges := by
have hin : (g.loopIn, i₁.natAdd 2) ∈ (Graph.loop g).edges := by
refine List.mem_append_left _ (List.mem_append_left _ (List.mem_append_right _ ?_))
exact List.mem_map_of_mem _ (List.mem_map_of_mem _ h₁)
-- the edge (2 ↑ʳ i₂) → out, reached through the third edge group
have hout : (i₂.natAdd 2, g.loopOut) ∈ (Graph.loop g).edges := by
refine List.mem_append_left _ (List.mem_append_right _ ?_) refine List.mem_append_left _ (List.mem_append_right _ ?_)
exact List.mem_map_of_mem _ (List.mem_map_of_mem _ h₂) exact List.mem_map_of_mem _ (List.mem_map_of_mem _ etr.exit_mem)
refine ⟨g.loopIn, List.mem_singleton_self _, g.loopOut, List.mem_singleton_self _, ?_⟩ refine ⟨g.loopIn, List.mem_singleton_self _, g.loopOut, List.mem_singleton_self _, ?_⟩
exact Trace.single (loop_nodes_at_in ▸ EvalBasicStmtOpt.none) ++< hin >++ exact (etr.beforeBody.appendTrace etr.trace.loop) ++< hout >++
tr.loop ++< hout >++ Trace.single (loop_nodes_at_out ▸ EvalBasicStmtOpt.none) 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. -/ /-- 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 : private lemma loop_edge_out_in :
@@ -141,16 +139,23 @@ private lemma loop_edge_out_in :
refine List.mem_append_right _ ?_ refine List.mem_append_right _ ?_
exact List.mem_cons_self _ _ exact List.mem_cons_self _ _
/-- Complete an iteration and follow the back edge before the remaining loop execution. -/
noncomputable def EndToEndTrace.beforeRest
(iteration : EndToEndTrace (Graph.loop g) ρ₁ ρ₂)
(rest : EndToEndTrace (Graph.loop g) ρ₂ ρ₃) :
Traceₗ (Graph.loop g) iteration.entry rest.entry ρ₁ ρ₂ := by
refine iteration.trace.addEdge ?_
have hout := iteration.exit_mem
have hin := rest.entry_mem
simp only [Graph.loop_inputs, Graph.loop_outputs, List.mem_singleton] at hin hout
simpa only [hin, hout] using (loop_edge_out_in (g := g))
/-- Two traces through a loop can be combined, since a loop can be executed any number of times. -/ /-- 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) ρ₁ ρ₂) noncomputable def EndToEndTrace.loop_concat (etr₁ : EndToEndTrace (Graph.loop g) ρ₁ ρ₂)
(etr₂ : EndToEndTrace (Graph.loop g) ρ₂ ρ₃) : (etr₂ : EndToEndTrace (Graph.loop g) ρ₂ ρ₃) :
EndToEndTrace (Graph.loop g) ρ₁ ρ₃ := by EndToEndTrace (Graph.loop g) ρ₁ ρ₃ := by
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁ exact ⟨etr₁.entry, etr₁.entry_mem, etr₂.exit, etr₂.exit_mem,
obtain ⟨j₁, k₁, j₂, k₂, tr₂⟩ := etr₂ etr₁.beforeRest etr₂ ++ etr₂.trace⟩
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 _,
tr₁ ++< loop_edge_out_in >++ tr₂⟩
/-- A loop can be executed zero times. -/ /-- A loop can be executed zero times. -/
noncomputable def EndToEndTrace.loop_empty {ρ : Env} : EndToEndTrace (Graph.loop g) ρ ρ := by noncomputable def EndToEndTrace.loop_empty {ρ : Env} : EndToEndTrace (Graph.loop g) ρ ρ := by
@@ -179,6 +184,14 @@ noncomputable def EndToEndTrace.wrap {g : Graph} {ρ₁ ρ₂ : Env}
(etr : EndToEndTrace g ρ₁ ρ₂) : EndToEndTrace (Graph.wrap g) ρ₁ ρ₂ := (etr : EndToEndTrace g ρ₁ ρ₂) : EndToEndTrace (Graph.wrap g) ρ₁ ρ₂ :=
(EndToEndTrace.singleton_nil ρ₁).concat (etr.concat (EndToEndTrace.singleton_nil ρ₂)) (EndToEndTrace.singleton_nil ρ₁).concat (etr.concat (EndToEndTrace.singleton_nil ρ₂))
/-- Reach the selected root entry through the program's empty wrapper node. -/
noncomputable def EndToEndTrace.beforeRoot {g : Graph} {ρ₁ ρ₂ : Env}
(root : EndToEndTrace g ρ₁ ρ₂) :
Traceₗ (Graph.wrap g) (Graph.wrapInput g)
(((GGraph.Embed.sequenceLeft g (Graph.singleton none)).trans
(GGraph.Embed.sequenceRight (Graph.singleton none) _)).f root.entry) ρ₁ ρ₁ :=
(EndToEndTrace.singleton_nil ρ₁).beforeRight (root.concat (EndToEndTrace.singleton_nil ρ₂))
/-- Key result: the control flow graph admits every execution that's made /-- 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 possible by a language's semantics. Thus, the CFG encodes _at least_ all
semantically-possible executions. Informally, we can conclude from this semantically-possible executions. Informally, we can conclude from this

View File

@@ -33,7 +33,7 @@ inductive Env.Mem : String × Value → Env → Prop
/-- Inference rules for evaluating an expression (`Spa.Expr`) in a given /-- Inference rules for evaluating an expression (`Spa.Expr`) in a given
environment. Pretty standard big-step expression evaluation. -/ environment. Pretty standard big-step expression evaluation. -/
inductive EvalExpr : Env → Expr → Value → Prop inductive EvalExpr : Env → Expr → Value → Prop
| num (ρ : Env) (n : ℕ) : EvalExpr ρ (.num n) (.int n) | num (ρ : Env) (z : ℤ) : EvalExpr ρ (.num z) (.int z)
| var (ρ : Env) (x : String) (v : Value) : | var (ρ : Env) (x : String) (v : Value) :
Env.Mem (x, v) ρ → EvalExpr ρ (.var x) v Env.Mem (x, v) ρ → EvalExpr ρ (.var x) v
| add (ρ : Env) (e₁ e₂ : Expr) (z₁ z₂ : ℤ) : | add (ρ : Env) (e₁ e₂ : Expr) (z₁ z₂ : ℤ) :

View File

@@ -1,18 +0,0 @@
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

View File

@@ -1,509 +0,0 @@
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

View File

@@ -1,104 +0,0 @@
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

View File

@@ -1,9 +0,0 @@
import Mathlib.Data.Nat.Notation
namespace Spa
structure RawId where
post : ℕ
deriving DecidableEq, Repr
end Spa

View File

@@ -1,29 +0,0 @@
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

View File

@@ -0,0 +1,114 @@
import Spa.Language.Properties
import Spa.Language.Equivalence
namespace Spa
open GGraph
/-- Recorded nodes contain instructions; empty CFG nodes are omitted from the history. -/
lemma Path.steps_nonempty {g : Graph} {a b : Configuration g} (p : Path g a b)
{d : g.Index} (hm : d ∈ p.steps) : g.nodes d ≠ none := by
induction p with
| nil => simp [Path.steps] at hm
| cons st p ih =>
rcases List.mem_append.mp hm with hs | hp
· cases st with
| edge => simp [Step.steps] at hs
| @execute i ρ σ h =>
cases hc : g.nodes i <;> aesop (add simp [Step.steps, hc])
· exact ih hp
private lemma optional_preserves_unwritten {ρ σ : Env} {obs : Option BasicStmt}
(h : EvalBasicStmtOpt ρ obs σ) (x : String)
(hn : ∀ rhs, obs ≠ some (.assign x rhs)) :
∀ v, Env.Mem (x, v) ρ ↔ Env.Mem (x, v) σ := by
cases h with
| none => exact fun _ => Iff.rfl
| some h =>
cases h with
| noop => exact fun _ => Iff.rfl
| assign y rhs w hv =>
have hxy : x ≠ y := by
rintro rfl
exact hn rhs rfl
intro v; simp [Env.mem_cons, hxy]
/-- A path whose executed nodes do not assign `x` preserves its binding. -/
lemma Path.preserves_unwritten {g : Graph} {a b : Configuration g} (p : Path g a b)
{x : String} (hn : ∀ d ∈ p.steps, ∀ rhs, g.nodes d ≠ some (.assign x rhs)) :
∀ v, Env.Mem (x, v) a.2 ↔ Env.Mem (x, v) b.2 := by
induction p with
| nil => exact fun _ => Iff.rfl
| cons st p ih =>
have ht := ih (fun d hm => hn d (List.mem_append_right _ hm))
suffices hs : ∀ v, Env.Mem (x, v) _ ↔ Env.Mem (x, v) _ from
fun v => (hs v).trans (ht v)
cases st with
| edge => exact fun _ => Iff.rfl
| execute h =>
apply optional_preserves_unwritten h x
intro rhs hc
exact hn _ (List.mem_append_left _ (by simp [Step.steps, hc])) rhs hc
lemma Step.steps_embed {g h : Graph} (e : Embed g h) {a b : Configuration g}
(s : Step g a b) :
(s.embed e).steps = s.steps.map e.f := by
cases s with
| edge => rfl
| @execute i ρ σ h =>
simp only [Step.embed, Step.steps, e.nodes_eq]
cases g.nodes i <;> rfl
lemma Path.steps_embed {g h : Graph} (e : Embed g h) {a b : Configuration g}
(p : Path g a b) :
(p.embed e).steps = p.steps.map e.f := by
induction p <;> aesop (add simp [Path.embed, Path.steps, Step.steps_embed])
/-- Every nonempty node in a loop belongs to its body. -/
lemma GGraph.loop_node_in_body {g : Graph} {i : (Graph.loop g).Index} {bs : BasicStmt}
(hc : (Graph.loop g).nodes i = some bs) : ∃ j, (Embed.loop g).f j = i := by
refine Fin.addCases ?_ ?_ i hc
· intro j hj
simp [Graph.loop, Fin.append_left] at hj
· intro j _; exact ⟨j, rfl⟩
/-- Variables at any CFG statement occur in its source statement. -/
lemma Stmt.cfg_node_vars {s : Stmt} {i : s.cfg.Index} {bs : BasicStmt}
(hc : s.cfg.nodes i = some bs) : bs.vars ⊆ s.vars := by
induction s with
| basic b =>
have : b = bs := Option.some.inj hc
subst bs; exact Finset.Subset.refl _
| andThen a b iha ihb =>
refine Fin.addCases ?_ ?_ i hc
· intro j hj; have hv := iha (by simpa [Stmt.cfg, Graph.sequence] using hj)
exact fun x hx => Finset.mem_union_left _ (hv hx)
· intro j hj; have hv := ihb (by simpa [Stmt.cfg, Graph.sequence] using hj)
exact fun x hx => Finset.mem_union_right _ (hv hx)
| ifElse cond a b iha ihb =>
refine Fin.addCases ?_ ?_ i hc
· intro j hj; have hv := iha (by simpa [Stmt.cfg, Graph.overlay] using hj)
exact fun x hx => Finset.mem_union_left _ (Finset.mem_union_right _ (hv hx))
· intro j hj; have hv := ihb (by simpa [Stmt.cfg, Graph.overlay] using hj)
exact fun x hx => Finset.mem_union_right _ (hv hx)
| whileLoop cond body ih =>
obtain ⟨j, rfl⟩ := GGraph.loop_node_in_body hc
have hv := ih (((Embed.loop body.cfg).nodes_eq j).symm.trans hc)
exact fun x hx => Finset.mem_union_right _ (hv hx)
lemma Program.code_vars {prog : Program} {i : prog.State} {bs : BasicStmt}
(hc : prog.code i = some bs) : ∀ x ∈ bs.vars, x ∈ prog.vars := by
have hroot : ∃ j, prog.rootStmt.cfg.nodes j = some bs := by
unfold Program.code Program.cfg Graph.wrap at hc
revert hc
refine Fin.addCases ?_ ?_ i
· intro j hj; simp [Graph.sequence, Graph.singleton] at hj
· intro j
refine Fin.addCases ?_ ?_ j
· intro k hk
exact ⟨k, by simpa [Graph.sequence] using hk⟩
· intro k hk; simp [Graph.sequence, Graph.singleton] at hk
obtain ⟨j, hj⟩ := hroot
intro x hx
simpa [Program.vars] using Stmt.cfg_node_vars hj hx
end Spa

View File

@@ -1,21 +1,22 @@
import Spa.Language.Semantics
import Spa.Language.Graphs import Spa.Language.Graphs
import Spa.Language.Program import Spa.Language.Program
import Spa.Language.Semantics
/-! /-!
# Program Traces # Program Traces
This module defines program traces tied to Control Flow Graphs, or CFGs 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 (see `Spa.GGraph` and `Spa.Graph`). These traces boil down to sequences of
basic-block executions (really, `Spa.BasicStmt` executions), each of which must 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 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 basic block by an edge. In this way, traces encode executions admitted
by the CFG. by the CFG.
While the regular `Trace` is just _any_ path through the graph, an `Path` interleaves execution and edge steps, with endpoints recording whether
`EndToEndTrace` is a path from the entry node to the exit node, denoting we are before or after a node. `Trace`, `Traceₗ`, and `Traceᵣ` are endpoint
full program execution. specializations of this one type. An `EndToEndTrace` runs from a graph input
to a graph output, denoting full program execution.
Properties about graphs and language semantics (especially, Properties about graphs and language semantics (especially,
the fact that the graph contains the proper basic block and edges the fact that the graph contains the proper basic block and edges
@@ -27,240 +28,236 @@ in `Spa/Language/Properties.lean`.
namespace Spa namespace Spa
/-- A partial trace through a graph `g`, starting right before /-- A node together with the phase of its execution. -/
the execution of the basic block at the first index, and inductive Position (α : Type) where
ending right after the execution of the basic block at the last index. -/ | before : α → Position α
inductive Trace (g : Graph) : g.Index → g.Index → Env → Env → Type | after : α → Position α
| single {ρ₁ ρ₂ : Env} {idx : g.Index} : deriving DecidableEq
EvalBasicStmtOpt ρ₁ (g.nodes idx) ρ₂ → Trace g idx idx ρ₁ ρ₂
| edge {ρ₁ ρ₂ ρ₃ : Env} {idx₁ idx₂ idx₃ : g.Index} :
EvalBasicStmtOpt ρ₁ (g.nodes idx₁) ρ₂ → (idx₁, idx₂) ∈ g.edges →
Trace g idx₂ idx₃ ρ₂ ρ₃ → Trace g idx₁ idx₃ ρ₁ ρ₃
/-! abbrev Configuration (g : Graph) := Position g.Index × Env
## Open Traces /-- Executing a node changes the environment; following an edge preserves it. -/
inductive Step (g : Graph) : Configuration g → Configuration g → Type where
| execute {i : g.Index} {ρ ρ' : Env}
(h : EvalBasicStmtOpt ρ (g.nodes i) ρ') :
Step g (.before i, ρ) (.after i, ρ')
| edge {i j : g.Index} {ρ : Env} (h : (i, j) ∈ g.edges) :
Step g (.after i, ρ) (.before j, ρ)
A normal `Trace` starts right before one state, and ends right after another. /-- A concrete CFG path, including executions of statement-less nodes. -/
This is convenient for inductively proving correctness / sufficience, but inductive Path (g : Graph) : Configuration g → Configuration g → Type where
awkward because 1) no empty traces exist and 2) concatenation requires an extra | nil {a} : Path g a a
edge. | cons {a b c} : Step g a b → Path g b c → Path g a c
However, when attempting an "empty" trace, two types are equally possible: namespace Path
traces that end _right before_ executing a state (`Traceₗ`) and
traces that begin _right after_ executing a state (`Traceᵣ`). They
are symmetric and can be concatenated with full traces on the left
and right, respectively. -/
/-- Left-open trace, representing execution that ends right before `idx₂`. -/ variable {g : Graph} {a b c d : Configuration g}
inductive Traceₗ (g : Graph) : g.Index → g.Index → Env → Env → Type where
| nil {idx : g.Index} {ρ : Env} : Traceₗ g idx idx ρ ρ
| cons {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} :
EvalBasicStmtOpt ρ₁ (g.nodes idx₁) ρ₂ →
(idx₁, idx₂) ∈ g.edges →
Traceₗ g idx₂ idx₃ ρ₂ ρ₃ → Traceₗ g idx₁ idx₃ ρ₁ ρ₃
def Traceₗ.single (g : Graph) (idx : g.Index) (ρ : Env) : Traceₗ g idx idx ρ ρ := .nil @[match_pattern] def single (s : Step g a b) : Path g a b := .cons s .nil
/-- Right-open trace, representing execution that starts right after `idx₁`. -/ def append {a b c : Configuration g} : Path g a b → Path g b c → Path g a c
inductive Traceᵣ (g : Graph) : g.Index → g.Index → Env → Env → Type where | .nil, q => q
| nil {idx : g.Index} {ρ : Env} : Traceᵣ g idx idx ρ ρ | .cons s p, q => .cons s (p.append q)
| cons {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} :
Traceᵣ g idx₁ idx₂ ρ₁ ρ₂ →
(idx₂, idx₃) ∈ g.edges →
EvalBasicStmtOpt ρ₂ (g.nodes idx₃) ρ₃ → Traceᵣ g idx₁ idx₃ ρ₁ ρ₃
def Traceᵣ.single (g : Graph) (idx : g.Index) (ρ : Env) : Traceᵣ g idx idx ρ ρ := .nil instance : HAppend (Path g a b) (Path g b c) (Path g a c) := ⟨append⟩
/-- Sequence two traces together. Since the endpoint of the first trace @[simp] lemma nil_append (p : Path g a b) : Path.nil.append p = p := rfl
is _after_ its last basic block's execution, and the beginning of
the next trace is _before_ its first basic block's execution, @[simp] lemma append_nil (p : Path g a b) : p.append Path.nil = p := by
there must be an edge to connect the two. -/ induction p <;> aesop (add simp append)
def Trace.concat {g : Graph} {idx₁ idx₂ idx₃ idx₄ : g.Index}
{ρ₁ ρ₂ ρ₃ : Env} (tr₁ : Trace g idx₁ idx₂ ρ₁ ρ₂) lemma append_assoc (p : Path g a b) (q : Path g b c) (r : Path g c d) :
(he : (idx₂, idx₃) ∈ g.edges) (tr₂ : Trace g idx₃ idx₄ ρ₂ ρ₃) : (p.append q).append r = p.append (q.append r) := by
Trace g idx₁ idx₄ ρ₁ ρ₃ := induction p <;> aesop (add simp append)
match tr₁ with
| single hbs => edge hbs he tr₂ end Path
| edge hbs he' tr₁' => edge hbs he' (tr₁'.concat he tr₂)
def GGraph.Embed.mapConfiguration {g h : Graph} (e : GGraph.Embed g h) :
Configuration g → Configuration h
| (.before i, ρ) => (.before (e.f i), ρ)
| (.after i, ρ) => (.after (e.f i), ρ)
lemma GGraph.Embed.mapConfiguration_trans {g h k : Graph}
(e : GGraph.Embed g h) (f : GGraph.Embed h k) (a : Configuration g) :
f.mapConfiguration (e.mapConfiguration a) = (e.trans f).mapConfiguration a := by
rcases a with ⟨_ | _, ρ⟩ <;> rfl
noncomputable def Step.embed {g h : Graph} (e : GGraph.Embed g h)
{a b : Configuration g} : Step g a b → Step h (e.mapConfiguration a) (e.mapConfiguration b)
| .execute h => .execute (_root_.cast (congrArg (EvalBasicStmtOpt _ · _) (e.nodes_eq _).symm) h)
| .edge h => .edge (e.edges_mem h)
noncomputable def Path.embed {g h : Graph} (e : GGraph.Embed g h)
{a b : Configuration g} : Path g a b → Path h (e.mapConfiguration a) (e.mapConfiguration b)
| .nil => .nil
| .cons s p => .cons (s.embed e) (p.embed e)
lemma Path.embed_append {g h : Graph} (e : GGraph.Embed g h)
{a b c : Configuration g} (p : Path g a b) (q : Path g b c) :
(p.append q).embed e = (p.embed e).append (q.embed e) := by
induction p <;> aesop (add simp [append, embed])
/-- Transport endpoints without changing the path. -/
def Path.cast {g : Graph} {a b a' b' : Configuration g}
(ha : a = a') (hb : b = b') (p : Path g a b) : Path g a' b' := ha ▸ hb ▸ p
lemma Path.embed_trans {g h k : Graph} (e : GGraph.Embed g h) (f : GGraph.Embed h k)
{a b : Configuration g} (p : Path g a b) :
((p.embed e).embed f).cast (e.mapConfiguration_trans f a)
(e.mapConfiguration_trans f b) = p.embed (e.trans f) := by
induction p with
| @nil a => rcases a with ⟨_ | _, ρ⟩ <;> rfl
| @cons a b c s p ih =>
rcases c with ⟨_ | _, ρ⟩ <;> cases s <;>
aesop (add simp [embed, Step.embed, cast, GGraph.Embed.mapConfiguration, cast_cast])
/-- A trace includes the executions of both endpoint nodes. -/
abbrev Trace (g : Graph) (i j : g.Index) (ρ ρ' : Env) :=
Path g (.before i, ρ) (.after j, ρ')
/-- A prefix ending before execution of its final node. -/
abbrev Traceₗ (g : Graph) (i j : g.Index) (ρ ρ' : Env) :=
Path g (.before i, ρ) (.before j, ρ')
/-- A suffix starting after execution of its initial node. -/
abbrev Traceᵣ (g : Graph) (i j : g.Index) (ρ ρ' : Env) :=
Path g (.after i, ρ) (.after j, ρ')
/-- Compatibility patterns for an execution and an execution-edge pair. -/
@[match_pattern] abbrev Trace.single {g : Graph} {ρ₁ ρ₂ : Env} {idx : g.Index}
(h : EvalBasicStmtOpt ρ₁ (g.nodes idx) ρ₂) : Trace g idx idx ρ₁ ρ₂ :=
.cons (.execute h) .nil
@[match_pattern] abbrev Trace.edge {g : Graph} {ρ₁ ρ₂ ρ₃ : Env}
{idx₁ idx₂ idx₃ : g.Index} (h : EvalBasicStmtOpt ρ₁ (g.nodes idx₁) ρ₂)
(he : (idx₁, idx₂) ∈ g.edges) (p : Trace g idx₂ idx₃ ρ₂ ρ₃) :
Trace g idx₁ idx₃ ρ₁ ρ₃ := Path.cons (.execute h) (.cons (.edge he) p)
@[match_pattern] abbrev Traceₗ.nil {g : Graph} {idx : g.Index} {ρ : Env} :
Traceₗ g idx idx ρ ρ := Path.nil
@[match_pattern] abbrev Traceₗ.cons {g : Graph} {ρ₁ ρ₂ ρ₃ : Env}
{idx₁ idx₂ idx₃ : g.Index} (h : EvalBasicStmtOpt ρ₁ (g.nodes idx₁) ρ₂)
(he : (idx₁, idx₂) ∈ g.edges) (p : Traceₗ g idx₂ idx₃ ρ₂ ρ₃) :
Traceₗ g idx₁ idx₃ ρ₁ ρ₃ := Path.cons (.execute h) (.cons (.edge he) p)
@[match_pattern] abbrev Traceᵣ.nil {g : Graph} {idx : g.Index} {ρ : Env} : Traceᵣ g idx idx ρ ρ := Path.nil
abbrev Traceᵣ.cons {g : Graph} {ρ₁ ρ₂ ρ₃ : Env} {idx₁ idx₂ idx₃ : g.Index}
(p : Traceᵣ g idx₁ idx₂ ρ₁ ρ₂) (he : (idx₂, idx₃) ∈ g.edges)
(h : EvalBasicStmtOpt ρ₂ (g.nodes idx₃) ρ₃) : Traceᵣ g idx₁ idx₃ ρ₁ ρ₃ :=
p.append (.cons (.edge he) (.single (.execute h)))
abbrev Traceₗ.single (g : Graph) (idx : g.Index) (ρ : Env) : Traceₗ g idx idx ρ ρ := .nil
abbrev Traceᵣ.single (g : Graph) (idx : g.Index) (ρ : Env) : Traceᵣ g idx idx ρ ρ := .nil
abbrev Trace.concat {g : Graph} {idx₁ idx₂ idx₃ idx₄ : g.Index} {ρ₁ ρ₂ ρ₃ : Env}
(p : Trace g idx₁ idx₂ ρ₁ ρ₂) (he : (idx₂, idx₃) ∈ g.edges)
(q : Trace g idx₃ idx₄ ρ₂ ρ₃) : Trace g idx₁ idx₄ ρ₁ ρ₃ :=
(p.append (.single (.edge he))).append q
scoped notation:65 tr₁:66 " ++< " he " >++ " tr₂:65 => Trace.concat tr₁ he tr₂ scoped notation:65 tr₁:66 " ++< " he " >++ " tr₂:65 => Trace.concat tr₁ he tr₂
def Trace.addEdge {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ : Env} : abbrev Trace.addEdge {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ : Env}
Trace g idx₁ idx₂ ρ₁ ρ₂ → (p : Trace g idx₁ idx₂ ρ₁ ρ₂) (he : (idx₂, idx₃) ∈ g.edges) :
(idx₂, idx₃) ∈ g.edges → Traceₗ g idx₁ idx₃ ρ₁ ρ₂ := p.append (.single (.edge he))
Traceₗ g idx₁ idx₃ ρ₁ ρ₂
| .single hnode, hedge => .cons hnode hedge .nil
| .edge hnode hedge' rest, hedge => .cons hnode hedge' (rest.addEdge hedge)
@[aesop simp] abbrev Traceₗ.append {g : Graph} {i j k : g.Index} {ρ₁ ρ₂ ρ₃ : Env}
def Traceₗ.append {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : (p : Traceₗ g i j ρ₁ ρ₂) (q : Traceₗ g j k ρ₂ ρ₃) : Traceₗ g i k ρ₁ ρ₃ :=
Traceₗ g idx₁ idx₂ ρ₁ ρ₂ → Traceₗ g idx₂ idx₃ ρ₂ ρ₃ → Path.append p q
Traceₗ g idx₁ idx₃ ρ₁ ρ₃
| .nil, rhs => rhs
| .cons hnode hedge rest, rhs => .cons hnode hedge (rest.append rhs)
@[simp] def traceₗ_append_nil {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ : Env} abbrev Traceₗ.appendTrace {g : Graph} {i j k : g.Index} {ρ₁ ρ₂ ρ₃ : Env}
{trₗ : Traceₗ g idx₁ idx₂ ρ₁ ρ₂} : trₗ.append Traceₗ.nil = trₗ := by (p : Traceₗ g i j ρ₁ ρ₂) (q : Trace g j k ρ₂ ρ₃) : Trace g i k ρ₁ ρ₃ :=
induction trₗ <;> aesop Path.append p q
def Traceₗ.appendTrace {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : abbrev Trace.appendRight {g : Graph} {i j k : g.Index} {ρ₁ ρ₂ ρ₃ : Env}
Traceₗ g idx₁ idx₂ ρ₁ ρ₂ → Trace g idx₂ idx₃ ρ₂ ρ₃ → (p : Trace g i j ρ₁ ρ₂) (q : Traceᵣ g j k ρ₂ ρ₃) : Trace g i k ρ₁ ρ₃ :=
Trace g idx₁ idx₃ ρ₁ ρ₃ Path.append p q
| .nil, rhs => rhs
| .cons hnode hedge rest, rhs => .edge hnode hedge (rest.appendTrace rhs)
def Traceₗ.appendStep {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : noncomputable abbrev Trace.embed {g h : Graph} (e : GGraph.Embed g h)
Traceₗ g idx₁ idx₂ ρ₁ ρ₂ → EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃ → {i j : g.Index} {ρ₁ ρ₂ : Env} (p : Trace g i j ρ₁ ρ₂) :
Trace g idx₁ idx₂ ρ₁ ρ₃ := fun trₗ hbs => trₗ.appendTrace (Trace.single hbs) Trace h (e.f i) (e.f j) ρ₁ ρ₂ := Path.embed e p
def Trace.appendRight {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : abbrev Traceₗ.appendStep {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ ρ₃ : Env}
Trace g idx₁ idx₂ ρ₁ ρ₂ → Traceᵣ g idx₂ idx₃ ρ₂ ρ₃ → (p : Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (h : EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃) :
Trace g idx₁ idx₃ ρ₁ ρ₃ Trace g idx₁ idx₂ ρ₁ ρ₃ := Path.append p (.single (.execute h))
| lhs, .nil => lhs
| lhs, .cons rest hedge hnode => Trace.concat (lhs.appendRight rest) hedge (.single hnode)
instance instHAppendTraceLTraceL {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : instance {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} :
HAppend (Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (Traceₗ g idx₂ idx₃ ρ₂ ρ₃) (Traceₗ g idx₁ idx₃ ρ₁ ρ₃) where HAppend (Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃)
hAppend := Traceₗ.append (Trace g idx₁ idx₂ ρ₁ ρ₃) := ⟨Traceₗ.appendStep⟩
instance instHAppendTraceLTrace {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : /-- The nonempty node executed by this step; edges and empty nodes are omitted. -/
HAppend (Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (Trace g idx₂ idx₃ ρ₂ ρ₃) (Trace g idx₁ idx₃ ρ₁ ρ₃) where def Step.steps {g : Graph} {a b : Configuration g} : Step g a b → List g.Index
hAppend := Traceₗ.appendTrace | .execute (i := i) _ =>
match g.nodes i with
| none => []
| some _ => [i]
| .edge _ => []
instance instHAppendTraceLStep {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} : /-- Executed nodes in chronological order; edges and empty nodes contribute nothing.
HAppend (Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃) (Trace g idx₁ idx₂ ρ₁ ρ₃) where The instruction at each node is given by `g.nodes`, rather than copied into the history. -/
hAppend := Traceₗ.appendStep def Path.steps {g : Graph} {a b : Configuration g} : Path g a b → List g.Index
instance instHAppendTraceTraceR {g : Graph} {idx₁ idx₂ idx₃ : g.Index} {ρ₁ ρ₂ ρ₃ : Env} :
HAppend (Trace g idx₁ idx₂ ρ₁ ρ₂) (Traceᵣ g idx₂ idx₃ ρ₂ ρ₃) (Trace g idx₁ idx₃ ρ₁ ρ₃) where
hAppend := Trace.appendRight
/-!
## Trace Steps
Analyses that care about *which statements executed* (e.g. reaching
definitions) need to project a trace down to its list of executed statements.
Defining that projection here, once, as a chronological mathlib `List` means
all the re-association facts about concatenating traces come for free from
`List.append_assoc` and friends, instead of being re-proven per analysis. -/
/-- The (index, statement) pairs executed by a single optional-statement step:
none if the node is empty, and the node's statement otherwise. -/
def EvalBasicStmtOpt.steps {α : Type*} (idx : α) {ρ₁ ρ₂ : Env} {obs : Option BasicStmt} :
EvalBasicStmtOpt ρ₁ obs ρ₂ → List (α × BasicStmt)
| .none => []
| .some (bs := bs) _ => [(idx, bs)]
/-- The statements executed by a left-open trace, in chronological order. -/
def Traceₗ.steps {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ : Env} :
Traceₗ g idx₁ idx₂ ρ₁ ρ₂ → List (g.Index × BasicStmt)
| .nil => [] | .nil => []
| .cons (idx₁ := idx) hnode _ rest => hnode.steps idx ++ rest.steps | .cons s p => s.steps ++ p.steps
/-- The statements executed by a trace, in chronological order. -/ abbrev Trace.steps {g : Graph} {i j : g.Index} {ρ₁ ρ₂ : Env}
def Trace.steps {g : Graph} {idx₁ idx₂ : g.Index} {ρ₁ ρ₂ : Env} : (p : Trace g i j ρ₁ ρ₂) : List g.Index := Path.steps p
Trace g idx₁ idx₂ ρ₁ ρ₂ → List (g.Index × BasicStmt) abbrev Traceₗ.steps {g : Graph} {i j : g.Index} {ρ₁ ρ₂ : Env}
| .single (idx := idx) hnode => hnode.steps idx (p : Traceₗ g i j ρ₁ ρ₂) : List g.Index := Path.steps p
| .edge (idx₁ := idx) hnode _ rest => hnode.steps idx ++ rest.steps abbrev Traceᵣ.steps {g : Graph} {i j : g.Index} {ρ₁ ρ₂ : Env}
(p : Traceᵣ g i j ρ₁ ρ₂) : List g.Index := Path.steps p
@[simp] lemma Traceₗ.steps_append {g : Graph} {idx₁ idx₂ idx₃ : g.Index} @[simp] lemma Path.steps_append {g : Graph} {a b c : Configuration g}
{ρ₁ ρ₂ ρ₃ : Env} (tr₁ : Traceₗ g idx₁ idx₂ ρ₁ ρ₂) (p : Path g a b) (q : Path g b c) :
(tr₂ : Traceₗ g idx₂ idx₃ ρ₂ ρ₃) : (p.append q).steps = p.steps ++ q.steps := by
(tr₁ ++ tr₂).steps = tr₁.steps ++ tr₂.steps := by induction p <;> aesop (add simp [append, steps, List.append_assoc])
show (tr₁.append tr₂).steps = _
induction tr₁ <;> simp [Traceₗ.append, Traceₗ.steps, *]
@[simp] lemma Traceₗ.steps_appendTrace {g : Graph} {idx₁ idx₂ idx₃ : g.Index}
{ρ₁ ρ₂ ρ₃ : Env} (tr₁ : Traceₗ g idx₁ idx₂ ρ₁ ρ₂)
(tr₂ : Trace g idx₂ idx₃ ρ₂ ρ₃) :
(tr₁ ++ tr₂).steps = tr₁.steps ++ tr₂.steps := by
show (tr₁.appendTrace tr₂).steps = _
induction tr₁ <;> simp [Traceₗ.appendTrace, Traceₗ.steps, Trace.steps, *]
@[simp] lemma Traceₗ.steps_appendStep {g : Graph} {idx₁ idx₂ : g.Index} @[simp] lemma Traceₗ.steps_appendStep {g : Graph} {idx₁ idx₂ : g.Index}
{ρ₁ ρ₂ ρ₃ : Env} (tr : Traceₗ g idx₁ idx₂ ρ₁ ρ₂) {ρ₁ ρ₂ ρ₃ : Env} (tr : Traceₗ g idx₁ idx₂ ρ₁ ρ₂)
(hbs : EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃) : (hbs : EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃) :
(tr ++ hbs).steps = tr.steps ++ hbs.steps idx₂ := (tr ++ hbs).steps = tr.steps ++ (Step.execute hbs).steps := by
Traceₗ.steps_appendTrace tr (Trace.single hbs) change Path.steps (Path.append tr (Path.single (.execute hbs))) = _
aesop (add simp [Trace.steps, Traceₗ.steps, Path.single, Path.steps, Step.steps])
@[simp] lemma Trace.steps_addEdge {g : Graph} {idx₁ idx₂ idx₃ : g.Index} @[simp] lemma Trace.steps_addEdge {g : Graph} {idx₁ idx₂ idx₃ : g.Index}
{ρ₁ ρ₂ : Env} (tr : Trace g idx₁ idx₂ ρ₁ ρ₂) {ρ₁ ρ₂ : Env} (tr : Trace g idx₁ idx₂ ρ₁ ρ₂) (he : (idx₂, idx₃) ∈ g.edges) :
(hedge : (idx₂, idx₃) ∈ g.edges) : (tr.addEdge he).steps = tr.steps := by
(tr.addEdge hedge).steps = tr.steps := by change Path.steps (Path.append tr (Path.single (.edge he))) = _
induction tr <;> simp [Trace.addEdge, Trace.steps, Traceₗ.steps, *] aesop (add simp [Trace.steps, Traceₗ.steps, Path.single, Path.steps, Step.steps])
@[simp] lemma Traceₗ.append_addEdge {g : Graph}
{idx₁ idx₂ idx₃ idx₄ : g.Index} {ρ₁ ρ₂ ρ₃ ρ₄ : Env}
(trₗ : Traceₗ g idx₁ idx₂ ρ₁ ρ₂)
(hnode : EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃)
(hedge : (idx₂, idx₃) ∈ g.edges)
(rest : Traceₗ g idx₃ idx₄ ρ₃ ρ₄) :
trₗ.append (Traceₗ.cons hnode hedge rest) =
(Trace.addEdge (trₗ.appendStep hnode) hedge).append rest := by
induction trₗ <;> simp [Traceₗ.append, Traceₗ.appendStep, Traceₗ.appendTrace, Trace.addEdge, *]
@[simp] lemma Traceₗ.appendTrace_addEdge {g : Graph}
{idx₁ idx₂ idx₃ idx₄ : g.Index} {ρ₁ ρ₂ ρ₃ ρ₄ : Env}
(trₗ : Traceₗ g idx₁ idx₂ ρ₁ ρ₂)
(hnode : EvalBasicStmtOpt ρ₂ (g.nodes idx₂) ρ₃)
(hedge : (idx₂, idx₃) ∈ g.edges)
(rest : Trace g idx₃ idx₄ ρ₃ ρ₄) :
trₗ.appendTrace (Trace.edge hnode hedge rest) =
(Trace.addEdge (trₗ.appendStep hnode) hedge).appendTrace rest := by
induction trₗ <;> simp [Traceₗ.appendTrace, Traceₗ.appendStep, Trace.addEdge, *]
/-- A beginning-to-end trace corresponding to the CFG `g`. -/ /-- A beginning-to-end trace corresponding to the CFG `g`. -/
inductive EndToEndTrace (g : Graph) (ρ₁ ρ₂ : Env) : Type structure EndToEndTrace (g : Graph) (ρ₁ ρ₂ : Env) : Type where
| intro (idx₁ : g.Index) (idx₁_mem : idx₁ ∈ g.inputs) intro ::
(idx₂ : g.Index) (idx₂_mem : idx₂ ∈ g.outputs) entry : g.Index
(trace : Trace g idx₁ idx₂ ρ₁ ρ₂) : EndToEndTrace g ρ₁ ρ₂ entry_mem : entry ∈ g.inputs
exit : g.Index
exit_mem : exit ∈ g.outputs
trace : Trace g entry exit ρ₁ ρ₂
inductive Reaches {prog : Program} : {s₁ s₂ : prog.State} → {ρ₁ ρ₂ : Env} → /-- Every trace splits into the prefix arriving at its last node and that node's execution. -/
Trace prog.cfg s₁ s₂ ρ₁ ρ₂ → def Trace.split {g : Graph} {i₁ i₂ : g.Index} {ρ₁ ρ₂ : Env} :
(s : prog.State) → (ρin ρout : Env) → Type Trace g i₁ i₂ ρ₁ ρ₂ → Σ ρ, Traceₗ g i₁ i₂ ρ₁ ρ × EvalBasicStmtOpt ρ (g.nodes i₂) ρ₂
| single_here {s₁ : prog.State} {ρ₁ ρ₂ : Env} | Trace.single h => ⟨_, .nil, h⟩
(hnode : EvalBasicStmtOpt ρ₁ (prog.code s₁) ρ₂) : | Trace.edge h he rest =>
Reaches (.single hnode) s₁ ρ₁ ρ₂ let ⟨ρ, pre, step⟩ := rest.split
| edge_here {s₁ s₂ s₃ : prog.State} {ρ₁ ρ₂ ρ₃ : Env} ⟨ρ, Traceₗ.cons h he pre, step⟩
(hnode : EvalBasicStmtOpt ρ₁ (prog.code s₁) ρ₂)
(hedge : (s₁, s₂) ∈ prog.cfg.edges) (rest : Trace prog.cfg s₂ s₃ ρ₂ ρ₃) :
Reaches (.edge hnode hedge rest) s₁ ρ₁ ρ₂
| edge_there {s₁ s₂ s₃ : prog.State} {ρ₁ ρ₂ ρ₃ : Env}
(hnode : EvalBasicStmtOpt ρ₁ (prog.code s₁) ρ₂)
(hedge : (s₁, s₂) ∈ prog.cfg.edges) (rest : Trace prog.cfg s₂ s₃ ρ₂ ρ₃)
{s : prog.State} {ρin ρout : Env} :
Reaches rest s ρin ρout →
Reaches (.edge hnode hedge rest) s ρin ρout
def Reaches.pre {prog : Program} {s₁ s₂ s: prog.State} @[simp] lemma Trace.split_append {g : Graph} {i₁ i₂ : g.Index} {ρ₁ ρ₂ : Env}
{ρ₁ ρ₂ ρin ρout : Env} {tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂} : (tr : Trace g i₁ i₂ ρ₁ ρ₂) : tr.split.2.1 ++ tr.split.2.2 = tr := by
(r : Reaches tr s ρin ρout) → Traceₗ prog.cfg s₁ s ρ₁ ρin match tr with
| .single_here _ => .nil | Trace.single h => rw [Trace.split.eq_1]; rfl
| .edge_here _ _ _ => .nil | Trace.edge h he rest =>
| .edge_there hnode hedge _ r => .cons hnode hedge r.pre have ih := Trace.split_append rest
rw [Trace.split.eq_2]
aesop (add simp [HAppend.hAppend, Traceₗ.appendStep, Path.append])
def Reaches.post {prog : Program} {s₁ s₂ s: prog.State} structure Reaches {prog : Program} (s : prog.State) (ρin ρout : Env) : Type where
{ρ₁ ρ₂ ρin ρout : Env} {tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂} : pre : Traceₗ prog.cfg prog.initialState s [] ρin
(r : Reaches tr s ρin ρout) → Trace prog.cfg s₁ s ρ₁ ρout step : EvalBasicStmtOpt ρin (prog.code s) ρout
| .single_here hnode => .single hnode
| .edge_here hnode _ _ => .single hnode
| .edge_there hnode hedge _ r => .edge hnode hedge r.post
def Reaches.first {prog : Program} {s₁ s₂ s: prog.State}
{ρ₁ ρ₂ ρin ρout : Env} {tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂} :
(r : Reaches tr s ρin ρout) → Σ ρ₁', Reaches tr s₁ ρ₁ ρ₁'
| .single_here hnode => ⟨_, .single_here hnode⟩
| .edge_here hnode hedge hrest => ⟨_, .edge_here hnode hedge hrest⟩
| .edge_there hnode hedge hrest tmp' => ⟨_, .edge_here hnode hedge hrest⟩
def Reaches.step {prog : Program} {s₁ s₂ s: prog.State}
{ρ₁ ρ₂ ρin ρout : Env} {tr : Trace prog.cfg s₁ s₂ ρ₁ ρ₂} :
(r : Reaches tr s ρin ρout) → EvalBasicStmtOpt ρin (prog.code s) ρout
| .single_here hnode => hnode
| .edge_here hnode hedge hrest => hnode
| .edge_there hnode hedge hrest tmp' => tmp'.step
/-- Forget the environment before the last evaluated state. -/
def Reaches.post {prog : Program} {s : prog.State} {ρin ρout : Env}
(r : Reaches s ρin ρout) : Trace prog.cfg prog.initialState s [] ρout :=
r.pre ++ r.step
end Spa end Spa

View File

@@ -0,0 +1,178 @@
import Spa.Language.Base
import Spa.Language.Program
import Spa.Lattice.FiniteMap
import Spa.Analysis.Constant
import Spa.Analysis.Forward
/-!
# Constant folding
Rewrites each assignment's right-hand side to a literal wherever the constant
analysis (`Spa/Analysis/Constant.lean`) pins its value down.
The traversal recurses over the plain `Stmt`, threading a `GGraph.Embed` of the
current subtree's CFG into the program's (`Program.rootEmbed`, then one
`Embed.trans` per descent). At an assignment, `Embed.singletonIndex` gives its
CFG state, and the facts to fold with are `Forward.joinForKey` at that state —
the join over predecessors, i.e. the values *entering* the node, which is what
the right-hand side reads. (`Forward.variablesAt` would be the values *leaving*
it, which already include this assignment's own effect.)
-/
namespace Spa
namespace ConstantTransform
open GGraph Forward
variable (prog : Program)
abbrev result := Forward.result ConstLattice prog
/-- Replace an expression by a literal when the analysis pins its value down,
recursing into its subexpressions otherwise. Whole subexpressions are tried
first, so `(x + 1) - x` folds outright when `x` is known, rather than only in its
leaves. -/
def foldExpr (vs : VariableValues ConstLattice prog) : Expr → Expr
| .num n => .num n
| .var k =>
match ConstAnalysis.eval prog (.var k) vs with
| .mk z => .num z
| _ => .var k
| .add a b =>
match ConstAnalysis.eval prog (.add a b) vs with
| .mk z => .num z
| _ => .add (foldExpr vs a) (foldExpr vs b)
| .sub a b =>
match ConstAnalysis.eval prog (.sub a b) vs with
| .mk z => .num z
| _ => .sub (foldExpr vs a) (foldExpr vs b)
/-- Constant-fold every assignment in a statement.
`sv` is the analysis result, taken as a parameter rather than read from `result`
at each node: it is a fixpoint computation, so recomputing it per assignment
would make folding quadratic in the analysis.
Guards of `ifElse`/`whileLoop` are deliberately left alone. `Stmt.cfg` gives a
conditional's guard no node at all (`ifElse` overlays the two branches), so there
is no state whose entry facts describe where it is evaluated. A `whileLoop`'s
guard does have a candidate — the loop header `GGraph.loopIn` — but tying the
guard's evaluation environment to that node needs a lemma that does not exist
yet, so folding it here would be an unproven soundness claim. -/
def foldStmt (sv : StateVariables ConstLattice prog) :
(s : Stmt) → Embed s.cfg prog.cfg → Stmt
| .basic .noop, _ => .basic .noop
| .basic (.assign k v), e =>
.basic (.assign k (foldExpr prog (joinForKey e.singletonIndex sv) v))
| .andThen s₁ s₂, e =>
.andThen (foldStmt sv s₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e))
(foldStmt sv s₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e))
| .ifElse cond s₁ s₂, e =>
.ifElse cond (foldStmt sv s₁ ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e))
(foldStmt sv s₂ ((Embed.overlayRight s₁.cfg s₂.cfg).trans e))
| .whileLoop cond body, e =>
.whileLoop cond (foldStmt sv body ((Embed.loop body.cfg).trans e))
/-- Constant-fold a whole program, running the analysis once. -/
def foldProgram : Stmt := foldStmt prog (result prog) prog.rootStmt prog.rootEmbed
/-! ## Correctness
Folding preserves meaning *provided the facts folded with actually hold of the
environment folded in*. That proviso is the whole content: `foldExpr` is sound
against any `vs` that over-approximates `ρ`, and it is the analysis engine's job
(`Forward.analyze_correct_at`) to supply such a `vs` at each program point. -/
variable {prog}
/-- If the analysis pins an expression to a constant and its facts hold of `ρ`,
then the expression really does evaluate to that constant. This is
`ValidExprEvaluator` specialised to the `.mk` case, where `interpConst` says
exactly `v = .int z`. -/
lemma eq_int_of_eval_mk {vs : VariableValues ConstLattice prog} {ρ : Env}
{e : Expr} {v : Value} {z : ℤ}
(hev : EvalExpr ρ e v) (hvs : ⟦vs⟧ ρ) (hz : ConstAnalysis.eval prog e vs = .mk z) :
v = .int z := by
have h := ValidExprEvaluator.valid (L := ConstLattice) (prog := prog) hev hvs
rw [show ExprEvaluator.eval e vs = ConstAnalysis.eval prog e vs from rfl, hz] at h
exact h
/-- **Expression folding is meaning-preserving.** Whenever `vs` over-approximates
`ρ`, the folded expression evaluates in `ρ` to whatever the original did. -/
theorem foldExpr_eval {vs : VariableValues ConstLattice prog} {ρ : Env} (hvs : ⟦vs⟧ ρ) :
∀ {e : Expr} {v : Value}, EvalExpr ρ e v → EvalExpr ρ (foldExpr prog vs e) v := by
intro e
induction e with
| num n => intro v hev; simpa [foldExpr] using hev
| var k =>
intro v hev
simp only [foldExpr]
split
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
· exact hev
| add a b iha ihb =>
intro v hev
simp only [foldExpr]
split
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
· cases hev with
| add _ _ z₁ z₂ h₁ h₂ => exact EvalExpr.add ρ _ _ z₁ z₂ (iha h₁) (ihb h₂)
| sub a b iha ihb =>
intro v hev
simp only [foldExpr]
split
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
· cases hev with
| sub _ _ z₁ z₂ h₁ h₂ => exact EvalExpr.sub ρ _ _ z₁ z₂ (iha h₁) (ihb h₂)
/-- Fold a source evaluation using its actual whole-program execution prefix. -/
noncomputable def foldStmt_eval (prog : Program) {s : Stmt} {ρ₀ ρ₁ : Env}
(h : EvalStmt ρ₀ s ρ₁) :
(e : Embed s.cfg prog.cfg) →
(pre : Traceₗ prog.cfg prog.initialState (e.f (Stmt.cfg_sufficient h).entry) [] ρ₀) →
EvalStmt ρ₀ (foldStmt prog (result prog) s e) ρ₁ := by
induction h with
| basic ρ₀ ρ₁ bs hbs =>
intro e pre
have hr : Reaches e.singletonIndex ρ₀ ρ₁ :=
⟨pre, by rw [Program.code, e.nodes_singletonIndex]; exact .some hbs⟩
cases hbs with
| noop => exact .basic _ _ _ (.noop _)
| assign x expr v hev =>
exact .basic _ _ _ (.assign _ _ _ _
(foldExpr_eval (ConstAnalysis.analyze_correct_at prog hr).1 hev))
| andThen ρ₀ ρ₁ ρ₂ s₁ s₂ h₁ h₂ ih₁ ih₂ =>
intro e pre
exact .andThen _ _ _ _ _
(ih₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e) pre)
(ih₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e)
(Path.append pre (Path.embed e
((Stmt.cfg_sufficient h₁).beforeRight (Stmt.cfg_sufficient h₂)))))
| ifTrue ρ₀ ρ₁ cond z s₁ s₂ hc hz h ih =>
intro e pre
exact .ifTrue _ _ _ _ _ _ hc hz
(ih ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e) pre)
| ifFalse ρ₀ ρ₁ cond s₁ s₂ hc h ih =>
intro e pre
exact .ifFalse _ _ _ _ _ hc
(ih ((Embed.overlayRight s₁.cfg s₂.cfg).trans e) pre)
| whileTrue ρ₀ ρ₁ ρ₂ cond z body hc hz hb hr ihb ihr =>
intro e pre
exact .whileTrue _ _ _ _ _ _ hc hz
(ihb ((Embed.loop body.cfg).trans e)
(Path.append pre (Path.embed e (Stmt.cfg_sufficient hb).beforeBody)))
(ihr e (Path.append pre (Path.embed e
((Stmt.cfg_sufficient hb).loop.beforeRest (Stmt.cfg_sufficient hr)))))
| whileFalse ρ cond body hc =>
intro e pre
exact .whileFalse _ _ _ hc
/-- Constant folding preserves every terminating source evaluation. -/
noncomputable def foldProgram_eval (prog : Program) {ρ : Env}
(h : EvalStmt [] prog.rootStmt ρ) : EvalStmt [] (foldProgram prog) ρ :=
foldStmt_eval prog h prog.rootEmbed (Stmt.cfg_sufficient h).beforeRoot
end ConstantTransform
end Spa

View File

@@ -1,92 +1,157 @@
import Spa.Analysis.Reaching import Spa.Analysis.Reaching
import Spa.Language.Tagged.Graphs import Spa.Language.Equivalence
/-! /-!
# Finding loop-invariant assignments (LICM groundwork) # Loop-invariant code motion
This wires the **reaching-definitions** analysis (`Spa/Analysis/Reaching.lean`) This wires the **reaching-definitions** analysis (`Spa/Analysis/Reaching.lean`)
to the **tagged AST** to *find* — not yet move — assignments inside a `while` to the AST to find assignments inside a `while` loop whose right-hand side
loop whose right-hand side depends only on definitions made *outside* the loop. depends only on definitions made outside the loop. `licmCandidates` reports
These are the candidates a later LICM pass could hoist. these assignments; `hoistProgram` moves eligible leading assignments.
The pipeline, for each assignment immediately enclosed by a loop: The traversal recurses over the plain `Stmt`, threading a `GGraph.Embed` of the
current subtree's CFG into the program's (`Program.rootEmbed`, then one
`Embed.trans` per descent). That embedding is what supplies program states:
1. locate its CFG state via the tagged-graph bridge (`Program.stateOfNodeId`); 1. at an assignment, its CFG state is `Embed.singletonIndex` — the subtree's CFG
2. read the reaching definitions at the assignment's *entry* is a `singleton`, so its sole node is the state, and `nodes_eq` proves it
(`joinForKey s result` — the join over predecessors, i.e. before the holds that very statement;
assignment itself runs); 2. read the reaching definitions at the assignment's *entry* (`joinForKey s
result` — the join over predecessors, i.e. before the assignment runs);
3. union the definition sets of the RHS variables; 3. union the definition sets of the RHS variables;
4. map each definition site back to its `RawId` (`Program.nodeIdOf`) and check 4. check no definition site lies in the loop body's CFG range. Every embedding is
it is **not** inside the loop body (structural `subtreeIds` membership). a constant index shift, so the body occupies the interval
`[off, off + size)` (`GGraph.Embed.mem_range_iff`) and the test is two
comparisons.
If every reaching definition of every RHS variable lies outside the loop, the 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 assignment is reported as loop-invariant. Hoisting additionally requires the
reaching definitions outside the loop"); transitive/iterated invariance and the assignment to lead the body, its destination to be absent from the guard, and
actual hoisting are out of scope here. no reassignment of that destination in the remaining body. The hoist is guarded
by the original condition, preserving zero-iteration behavior.
`LicmTransformation.hoistProgram_eval` in `Spa/Transformation/Licm/Correctness.lean`
proves preservation of terminating executions and observable final bindings.
Transitive invariance and motion of non-leading assignments are not implemented.
-/ -/
namespace Spa namespace Spa
namespace LicmTransformation namespace LicmTransformation
open Forward open Forward GGraph
/-- The CFG footprint of an enclosing loop: its entry node (for reporting) and
the index interval its body occupies. -/
structure Enclosing (prog : Program) where
/-- The loop's entry node, i.e. `GGraph.loopIn` embedded into the program. -/
loopState : prog.State
/-- Start of the body's index range. -/
bodyOff : ℕ
/-- Length of the body's index range. -/
bodySize : ℕ
/-- Is this definition site inside the loop body's CFG range? -/
def Enclosing.covers {prog : Program} (l : Enclosing prog) (d : prog.State) : Bool :=
decide (l.bodyOff ≤ d.val ∧ d.val < l.bodyOff + l.bodySize)
/-- An assignment found inside a loop, paired with the data needed to test its /-- An assignment found inside a loop, paired with the data needed to test its
invariance against that (immediately enclosing) loop. -/ invariance against that (immediately enclosing) loop. -/
structure Candidate (prog : Program) where structure Candidate (prog : Program) where
/-- The enclosing `whileLoop`'s tag (for reporting). -/ /-- The enclosing loop. -/
loopId : prog.NodeId encl : Enclosing prog
/-- Every node id inside the loop body (the "is-child-of-loop" set). -/ /-- The assignment's CFG state. -/
bodyIds : List prog.NodeId assignState : prog.State
/-- The assignment `BasicStmt`'s tag — what labels its CFG node. -/
assignId : prog.NodeId
/-- The variables read by the assignment's RHS. -/ /-- The variables read by the assignment's RHS. -/
rhsVars : List String rhsVars : List String
/-- Collect every assignment together with its *immediately enclosing* loop. /-- Collect every assignment together with its *immediately enclosing* loop.
`enclosing` carries the current loop's tag and body id-set, or `none` outside any `enc` is `none` outside any loop, in which case assignments are skipped — only
loop (in which case assignments are skipped — only in-loop assignments are in-loop assignments are candidates. -/
candidates). -/ def collectCandidates (prog : Program) (enc : Option (Enclosing prog)) :
def collectCandidates (prog : Program) (enc : Option (prog.NodeId × List prog.NodeId)) : (s : Stmt) → Embed s.cfg prog.cfg → List (Candidate prog)
Stmt.Tagged prog.NodeId → List (Candidate prog) | .basic bs, e =>
| .basic _ bs =>
match bs, enc with match bs, enc with
| .assign t _ e, some (loopId, bodyIds) => | .assign _ ex, some l =>
[{ loopId := loopId, bodyIds := bodyIds, assignId := t, [{ encl := l, assignState := e.singletonIndex,
rhsVars := e.erase.vars.sort (· ≤ ·) }] rhsVars := ex.vars.sort (· ≤ ·) }]
| _, _ => [] | _, _ => []
| .andThen _ a b => collectCandidates prog enc a ++ collectCandidates prog enc b | .andThen s₁ s₂, e =>
| .ifElse _ _ a b => collectCandidates prog enc a ++ collectCandidates prog enc b collectCandidates prog enc s₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e) ++
| .whileLoop loopT _ body => collectCandidates prog enc s₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e)
collectCandidates prog (some (loopT, body.subtreeIds)) body | .ifElse _ s₁ s₂, e =>
collectCandidates prog enc s₁ ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e) ++
collectCandidates prog enc s₂ ((Embed.overlayRight s₁.cfg s₂.cfg).trans e)
| .whileLoop _ body, e =>
let be := (Embed.loop body.cfg).trans e
collectCandidates prog
(some { loopState := e.f body.cfg.loopIn, bodyOff := be.off,
bodySize := body.cfg.size }) body be
/-- Read the definition set assigned to variable `k`, or `⊥` if absent. -/ /-- Read the definition set assigned to variable `k`, or `⊥` if absent. -/
def lookupDef (prog : Program) (vs : VariableValues (DefSet prog) prog) def lookupDef (prog : Program) (vs : VariableValues (DefSet prog) prog)
(k : String) : DefSet prog := (k : String) : DefSet prog :=
if h : FiniteMap.MemKey k vs then (FiniteMap.locate h).1 else ⊥ if h : FiniteMap.MemKey k vs then (FiniteMap.locate h).1 else ⊥
/-- The AST node ids marked as definition sites in a `DefSet`. With the
`Finset`-of-AST-ids lattice these are just the elements of the set. -/
def defSites (prog : Program) (d : DefSet prog) : List prog.NodeId :=
(List.finRange prog.size).filter (fun i => decide (i ∈ d))
/-- Is the candidate assignment loop-invariant: do all reaching definitions of /-- 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 its RHS variables lie outside the loop body? -/
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 := def isInvariant (prog : Program) (c : Candidate prog) : Bool :=
match prog.stateOfNodeId c.assignId with let entry := joinForKey c.assignState (result (DefSet prog) prog)
| none => false
| some s =>
let entry := joinForKey s (result (DefSet prog) prog)
let combined : DefSet prog := let combined : DefSet prog :=
c.rhsVars.foldl (fun acc k => acc ⊔ lookupDef prog entry k) ⊥ c.rhsVars.foldl (fun acc k => acc ⊔ lookupDef prog entry k) ⊥
(defSites prog combined).all (fun nid => ! decide (nid ∈ c.bodyIds)) -- `Finset.toList` is noncomputable; the decidable bounded-∀ folds over the
-- underlying multiset and keeps `lake exe` working.
decide (∀ d ∈ combined, c.encl.covers d = false)
/-- The loop-invariant assignments of `prog`, as `(loopId, assignId)` pairs. -/ /-- The loop-invariant assignments of `prog`, as `(loop, assignment)` state pairs. -/
def licmCandidates (prog : Program) : List (prog.NodeId × prog.NodeId) := def licmCandidates (prog : Program) : List (prog.State × prog.State) :=
(collectCandidates prog none prog.taggedFin).filterMap (fun c => (collectCandidates prog none prog.rootStmt prog.rootEmbed).filterMap (fun c =>
if isInvariant prog c then some (c.loopId, c.assignId) else none) if isInvariant prog c then some (c.encl.loopState, c.assignState) else none)
/-- Candidate for the leading assignment of a loop body. -/
def headCandidate (prog : Program) (cond : Expr) (x : String) (rhs : Expr) (tail : Stmt)
(e : Embed (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg prog.cfg) :
Candidate prog :=
let body := Stmt.andThen (.basic (.assign x rhs)) tail
let be := (Embed.loop body.cfg).trans e
let ae := (Embed.sequenceLeft (Stmt.basic (.assign x rhs)).cfg tail.cfg).trans be
{ encl := { loopState := e.f body.cfg.loopIn, bodyOff := be.off, bodySize := body.cfg.size },
assignState := ae.singletonIndex, rhsVars := rhs.vars.sort (· ≤ ·) }
/-- Guard the hoist so that a zero-iteration loop never evaluates the RHS. -/
def hoistHead (cond : Expr) (x : String) (rhs : Expr) (tail : Stmt) : Stmt :=
.ifElse cond (.andThen (.basic (.assign x rhs)) (.whileLoop cond tail)) (.basic .noop)
/-- Hoist a leading invariant assignment when its destination is neither
reassigned in the remaining body nor read by the guard. -/
def hoistLoop (prog : Program) (cond : Expr) (body : Stmt)
(e : Embed (Stmt.whileLoop cond body).cfg prog.cfg) : Option Stmt :=
match body with
| .andThen (.basic (.assign x rhs)) tail =>
if isInvariant prog (headCandidate prog cond x rhs tail e) &&
decide (x ∉ tail.writes ∧ x ∉ cond.vars) then
some (hoistHead cond x rhs tail)
else none
| _ => none
/-- Apply guarded leading-assignment LICM throughout the source tree. When a
loop is hoisted, keep its remaining body intact; further passes can reanalyze it. -/
def hoistStmt (prog : Program) : (s : Stmt) → Embed s.cfg prog.cfg → Stmt
| .basic bs, _ => .basic bs
| .andThen a b, e =>
.andThen (hoistStmt prog a ((Embed.sequenceLeft a.cfg b.cfg).trans e))
(hoistStmt prog b ((Embed.sequenceRight a.cfg b.cfg).trans e))
| .ifElse cond a b, e =>
.ifElse cond (hoistStmt prog a ((Embed.overlayLeft a.cfg b.cfg).trans e))
(hoistStmt prog b ((Embed.overlayRight a.cfg b.cfg).trans e))
| .whileLoop cond body, e =>
match hoistLoop prog cond body e with
| some moved => moved
| none => .whileLoop cond (hoistStmt prog body ((Embed.loop body.cfg).trans e))
/-- Run reaching definitions on the source program and perform guarded LICM. -/
def hoistProgram (prog : Program) : Stmt := hoistStmt prog prog.rootStmt prog.rootEmbed
/-- A human-readable report of the loop-invariant assignments. -/ /-- A human-readable report of the loop-invariant assignments. -/
def output (prog : Program) : String := def output (prog : Program) : String :=

View File

@@ -0,0 +1,259 @@
import Spa.Transformation.Licm
import Spa.Analysis.Reaching.Paths
namespace Spa
namespace LicmTransformation
open GGraph Forward ReachingAnalysis
private lemma mem_union_fold {α β : Type} [DecidableEq β] (f : α → Finset β)
(xs : List α) (acc : Finset β) (d : β) :
d ∈ xs.foldl (fun a x => a ∪ f x) acc ↔ d ∈ acc ∨ ∃ x ∈ xs, d ∈ f x := by
induction xs generalizing acc with
| nil => simp
| cons x xs ih => simp [List.foldl, ih, or_assoc, or_left_comm, or_comm]
/-- The executable invariant test excludes each actual reaching definition. -/
lemma isInvariant_sound_at {prog : Program} {c : Candidate prog} {ρ ρ' : Env}
{x : String} {d : prog.State}
(hinv : isInvariant prog c = true) (hx : x ∈ c.rhsVars) (hxp : x ∈ prog.vars)
(hr : Reaches c.assignState ρ ρ')
(hl : LastAssign prog x (runOfPath prog hr.pre) d) : c.encl.covers d = false := by
let entry := joinForKey c.assignState (result (DefSet prog) prog)
have hk : FiniteMap.MemKey x entry := hxp
have hd : d ∈ lookupDef prog entry x := by
have hs := (ReachingAnalysis.analyze_correct_at prog hr).1
have hm := (FiniteMap.locate hk).2
have hd := hs x (FiniteMap.locate hk).1 hm d hl
simpa [lookupDef, hk] using hd
have hall : ∀ d ∈ c.rhsVars.foldl (fun acc k => acc ∪ lookupDef prog entry k) ∅,
c.encl.covers d = false := by simpa [isInvariant, entry] using hinv
exact hall d ((mem_union_fold _ _ _ _).mpr (Or.inr ⟨x, hx, hd⟩))
/-- A path inside the loop can execute statements only within its body range. -/
lemma loop_steps_covered {prog : Program} {cond : Expr} {x : String} {rhs : Expr} {tail : Stmt}
(e : Embed (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg prog.cfg)
{a b : Configuration (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg}
(seg : Path _ a b) {d : prog.State}
(hm : d ∈ (seg.embed e).steps) :
(headCandidate prog cond x rhs tail e).encl.covers d = true := by
rw [Path.steps_embed] at hm
obtain ⟨j, hj, rfl⟩ := List.mem_map.mp hm
obtain ⟨bs, hcode⟩ := Option.ne_none_iff_exists'.mp (seg.steps_nonempty hj)
obtain ⟨i, hi⟩ := GGraph.loop_node_in_body hcode
apply decide_eq_true
apply (Embed.mem_range_iff ((Embed.loop _).trans e) _).mp
exact ⟨i, congrArg e.f hi⟩
/-- The analysis and the actual intervening path together establish RHS
stability. This is the bridge from static sites to unchanged runtime values. -/
lemma head_rhs_agrees {prog : Program} {cond : Expr} {x : String} {rhs : Expr} {tail : Stmt}
(e : Embed (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg prog.cfg)
{i j : (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg.Index}
{ρ₀ ρ₁ ρ₂ : Env}
(pre : Traceₗ prog.cfg prog.initialState (e.f i) [] ρ₀)
(seg : Traceₗ _ i j ρ₀ ρ₁)
(hj : e.f j = (headCandidate prog cond x rhs tail e).assignState)
(step : EvalBasicStmt ρ₁ (.assign x rhs) ρ₂)
(hinv : isInvariant prog (headCandidate prog cond x rhs tail e) = true) :
Env.AgreeOn rhs.vars ρ₀ ρ₁ := by
have hj' : j = ((Embed.sequenceLeft (Stmt.basic (.assign x rhs)).cfg tail.cfg).trans
(Embed.loop (Stmt.andThen (.basic (.assign x rhs)) tail).cfg)).singletonIndex := e.f_inj hj
subst j
let c := headCandidate prog cond x rhs tail e
have hcode : prog.code c.assignState = some (.assign x rhs) :=
Embed.nodes_singletonIndex
((Embed.sequenceLeft (Stmt.basic (.assign x rhs)).cfg tail.cfg).trans
((Embed.loop (Stmt.andThen (.basic (.assign x rhs)) tail).cfg).trans e))
let reach : Reaches c.assignState ρ₁ ρ₂ :=
⟨pre.append (seg.embed e), hcode ▸ .some step⟩
intro y hy
apply ReachingAnalysis.Path.preserves_of_lastAssign_outside pre (seg.embed e)
{d | c.encl.covers d = true}
· intro d hm; exact loop_steps_covered e seg hm
· intro d hl hd
have hl' : LastAssign prog y (runOfPath prog reach.pre) d := hl
have hf := isInvariant_sound_at hinv (by simpa [c, headCandidate] using hy)
(Program.code_vars hcode y (Finset.mem_union_right _ hy)) reach hl'
exact Bool.noConfusion (hd.symm.trans hf)
private lemma loop_entry {cond : Expr} {body : Stmt} {ρ σ : Env}
(h : EvalStmt ρ (.whileLoop cond body) σ) :
(Stmt.cfg_sufficient h).entry = body.cfg.loopIn := by
cases h <;> rfl
/-- Remove all subsequent executions of the leading assignment. The accumulated
source path, rather than transformed histories, supplies the analysis facts. -/
private noncomputable def removeHead_eval (prog : Program)
{cond : Expr} {x : String} {rhs : Expr} {tail : Stmt}
{ρ ρ' : Env} (h : EvalStmt ρ (.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)) ρ')
(e : Embed (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg prog.cfg)
(hinv : isInvariant prog (headCandidate prog cond x rhs tail e) = true)
(hwrite : x ∉ tail.writes)
{base : Env} {v : Value}
(pre : Traceₗ prog.cfg prog.initialState
(e.f (Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn) [] base)
(hv : EvalExpr base rhs v)
(seg : Traceₗ (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn base ρ)
{σ : Env} (heq : Env.Equiv ρ σ) (hx : Env.Mem (x, v) σ) :
Σ σ', {_h : EvalStmt σ (.whileLoop cond tail) σ' // Env.Equiv ρ' σ'} := by
generalize hs : Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail) = s at h
induction h generalizing σ with
| basic => cases hs
| andThen => cases hs
| ifTrue => cases hs
| ifFalse => cases hs
| whileFalse ρ cond' body hc =>
cases hs
exact ⟨σ, .whileFalse _ _ _ (hc.congr_env (fun y _ => heq y)), heq⟩
| whileTrue ρ₀ ρ₁ ρ₂ cond' z body hc hz hb hr ihb ihr =>
cases hs
cases hb with
| andThen _ ρa _ _ _ ha ht =>
cases ha with
| basic _ _ _ ha =>
cases ha with
| assign _ _ w hw =>
let hb := EvalStmt.andThen _ _ _ _ _ (.basic _ _ _ (.assign _ _ _ _ hw)) ht
let toAssign := seg.append (Stmt.cfg_sufficient hb).beforeBody
have hagree := head_rhs_agrees e pre toAssign rfl (.assign _ _ _ _ hw) hinv
have hwv : w = v := hw.deterministic (hv.congr_env hagree)
subst w
have heqa : Env.Equiv ((x, v) :: ρ₀) σ :=
(heq.cons x v).trans (Env.cons_equiv_of_mem hx)
obtain ⟨σ₁, ht', heq₁⟩ := ht.congr_env heqa
have hx₁ : Env.Mem (x, v) σ₁ := (ht'.preserves_unwritten hwrite v).mp hx
let next := seg.append ((Stmt.cfg_sufficient hb).loop.beforeRest (Stmt.cfg_sufficient hr))
have next' : Traceₗ (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn base ρ₁ := by
simpa only [loop_entry hr] using next
obtain ⟨σ₂, hr', heq₂⟩ := ihr next' heq₁ hx₁ rfl
exact ⟨σ₂, .whileTrue _ _ _ _ _ _ (hc.congr_env (fun y _ => heq y)) hz ht' hr', heq₂⟩
/-- Guarded hoisting preserves every terminating execution of an eligible loop,
including its current bindings and definedness. -/
noncomputable def hoistHead_eval (prog : Program)
{cond : Expr} {x : String} {rhs : Expr} {tail : Stmt} {ρ ρ' σ : Env}
(h : EvalStmt ρ (.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)) ρ')
(e : Embed (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg prog.cfg)
(pre : Traceₗ prog.cfg prog.initialState (e.f (Stmt.cfg_sufficient h).entry) [] ρ)
(hinv : isInvariant prog (headCandidate prog cond x rhs tail e) = true)
(hwrite : x ∉ tail.writes) (hguard : x ∉ cond.vars)
(heq : Env.Equiv ρ σ) :
Σ σ', {_h : EvalStmt σ (hoistHead cond x rhs tail) σ' // Env.Equiv ρ' σ'} := by
cases h with
| whileFalse _ _ _ hc =>
exact ⟨σ, .ifFalse _ _ _ _ _ (hc.congr_env (fun y _ => heq y))
(.basic _ _ _ (.noop _)), heq⟩
| whileTrue ρ₀ ρ₁ ρ₂ _ z _ hc hz hb hr =>
cases hb with
| andThen _ ρa _ _ _ ha ht =>
cases ha with
| basic _ _ _ ha =>
cases ha with
| assign _ _ v hv =>
let hb := EvalStmt.andThen _ _ _ _ _ (.basic _ _ _ (.assign _ _ _ _ hv)) ht
have hc' := hc.congr_env (fun y _ => heq y)
have hc'' : EvalExpr ((x, v) :: σ) cond (.int z) := by
apply hc'.congr_env
intro y hy w
have hne : y ≠ x := by rintro rfl; exact hguard hy
simp [Env.mem_cons, hne]
obtain ⟨σ₁, ht', heq₁⟩ := ht.congr_env (heq.cons x v)
have hx₁ : Env.Mem (x, v) σ₁ :=
(ht'.preserves_unwritten hwrite v).mp (.here _ _ _)
let next := (Stmt.cfg_sufficient hb).loop.beforeRest (Stmt.cfg_sufficient hr)
have next' : Traceₗ (Stmt.whileLoop cond (.andThen (.basic (.assign x rhs)) tail)).cfg
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn
(Stmt.andThen (.basic (.assign x rhs)) tail).cfg.loopIn ρ ρ₁ := by
simpa only [loop_entry hr] using next
obtain ⟨σ₂, hr', heq₂⟩ := removeHead_eval prog hr e hinv hwrite pre hv next' heq₁ hx₁
exact ⟨σ₂, .ifTrue _ _ _ _ _ _ hc' hz
(.andThen _ _ _ _ _
(.basic _ _ _ (.assign _ _ _ _ (hv.congr_env (fun y _ => heq y))))
(.whileTrue _ _ _ _ _ _ hc'' hz ht' hr')), heq₂⟩
private noncomputable def hoistLoop_eval (prog : Program)
{cond : Expr} {body moved : Stmt} {ρ ρ' σ : Env}
(h : EvalStmt ρ (.whileLoop cond body) ρ')
(e : Embed (Stmt.whileLoop cond body).cfg prog.cfg)
(pre : Traceₗ prog.cfg prog.initialState (e.f (Stmt.cfg_sufficient h).entry) [] ρ)
(hm : hoistLoop prog cond body e = some moved) (heq : Env.Equiv ρ σ) :
Σ σ', {_h : EvalStmt σ moved σ' // Env.Equiv ρ' σ'} := by
unfold hoistLoop at hm
split at hm
· rename_i x rhs tail e
split at hm
· rename_i hc
simp only [Bool.and_eq_true, decide_eq_true_eq] at hc
cases hm
exact hoistHead_eval prog h e pre hc.1 hc.2.1 hc.2.2 heq
· cases hm
· cases hm
/-- Transform a source evaluation using its source CFG prefix. Recursive calls
can start in any environment with the same current bindings. -/
noncomputable def hoistStmt_eval (prog : Program) {s : Stmt} {ρ ρ' : Env}
(h : EvalStmt ρ s ρ') :
(e : Embed s.cfg prog.cfg) →
(pre : Traceₗ prog.cfg prog.initialState (e.f (Stmt.cfg_sufficient h).entry) [] ρ) →
∀ {σ}, Env.Equiv ρ σ →
Σ σ', {_h : EvalStmt σ (hoistStmt prog s e) σ' // Env.Equiv ρ' σ'} := by
induction h with
| basic ρ₀ ρ₁ bs hb =>
intro e pre σ heq
exact (EvalStmt.basic _ _ _ hb).congr_env heq
| andThen ρ₀ ρ₁ ρ₂ a b ha hb iha ihb =>
intro e pre σ heq
obtain ⟨σ₁, ha', heq₁⟩ := iha ((Embed.sequenceLeft a.cfg b.cfg).trans e) pre heq
obtain ⟨σ₂, hb', heq₂⟩ := ihb ((Embed.sequenceRight a.cfg b.cfg).trans e)
(pre.append (Path.embed e ((Stmt.cfg_sufficient ha).beforeRight (Stmt.cfg_sufficient hb)))) heq₁
exact ⟨σ₂, .andThen _ _ _ _ _ ha' hb', heq₂⟩
| ifTrue ρ₀ ρ₁ cond z a b hc hz h ih =>
intro e pre σ heq
obtain ⟨σ', h', heq'⟩ := ih ((Embed.overlayLeft a.cfg b.cfg).trans e) pre heq
exact ⟨σ', .ifTrue _ _ _ _ _ _ (hc.congr_env (fun x _ => heq x)) hz h', heq'⟩
| ifFalse ρ₀ ρ₁ cond a b hc h ih =>
intro e pre σ heq
obtain ⟨σ', h', heq'⟩ := ih ((Embed.overlayRight a.cfg b.cfg).trans e) pre heq
exact ⟨σ', .ifFalse _ _ _ _ _ (hc.congr_env (fun x _ => heq x)) h', heq'⟩
| whileTrue ρ₀ ρ₁ ρ₂ cond z body hc hz hb hr ihb ihr =>
intro e pre σ heq
cases hm : hoistLoop prog cond body e with
| some moved =>
simp only [hoistStmt]
rw [hm]
exact hoistLoop_eval prog (.whileTrue _ _ _ _ _ _ hc hz hb hr) e pre hm heq
| none =>
obtain ⟨σ₁, hb', heq₁⟩ := ihb ((Embed.loop body.cfg).trans e)
(pre.append (Path.embed e (Stmt.cfg_sufficient hb).beforeBody)) heq
obtain ⟨σ₂, hr', heq₂⟩ := ihr e
(pre.append (Path.embed e
((Stmt.cfg_sufficient hb).loop.beforeRest (Stmt.cfg_sufficient hr)))) heq₁
simp only [hoistStmt] at hr' ⊢
rw [hm] at hr' ⊢
exact ⟨σ₂, .whileTrue _ _ _ _ _ _ (hc.congr_env (fun x _ => heq x)) hz hb' hr', heq₂⟩
| whileFalse ρ cond body hc =>
intro e pre σ heq
cases hm : hoistLoop prog cond body e with
| some moved =>
simp only [hoistStmt]
rw [hm]
exact hoistLoop_eval prog (.whileFalse _ _ _ hc) e pre hm heq
| none =>
simp only [hoistStmt]
rw [hm]
exact ⟨σ, .whileFalse _ _ _ (hc.congr_env (fun x _ => heq x)), heq⟩
/-- LICM preserves every terminating source execution and all observable final
bindings. The source analysis is computed by `hoistProgram`; no soundness or
invariance premise is required of callers. -/
noncomputable def hoistProgram_eval (prog : Program) {ρ : Env}
(h : EvalStmt [] prog.rootStmt ρ) :
Σ σ, {_h : EvalStmt [] (hoistProgram prog) σ // Env.Equiv ρ σ} :=
hoistStmt_eval prog h prog.rootEmbed (Stmt.cfg_sufficient h).beforeRoot (Env.Equiv.refl [])
end LicmTransformation
end Spa