Prove correctness of constant folding

This commit is contained in:
2026-10-04 10:27:37 -05:00
parent d2b6bf5af7
commit 234d17394e
2 changed files with 93 additions and 23 deletions

View File

@@ -71,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₂),
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
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
@@ -80,13 +89,10 @@ noncomputable def EndToEndTrace.overlay_right (etr : EndToEndTrace g₂ ρ₁ ρ
with a trace in another graph. -/
noncomputable def EndToEndTrace.concat {ρ₃ : Env} (etr₁ : EndToEndTrace g₁ ρ₁ ρ₂)
(etr₂ : EndToEndTrace g₂ ρ₂ ρ₃) : EndToEndTrace (g₁ ⤳ g₂) ρ₁ ρ₃ := by
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁
obtain ⟨j₁, k₁, j₂, k₂, tr₂⟩ := etr₂
refine ⟨i₁.castAdd g₂.size, List.mem_map_of_mem _ h₁,
j₂.natAdd g₁.size, List.mem_map_of_mem _ k₂,
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₁⟩)
exact ⟨etr₁.entry.castAdd g₂.size, List.mem_map_of_mem _ etr₁.entry_mem,
etr₂.exit.natAdd g₁.size, List.mem_map_of_mem _ etr₂.exit_mem,
(etr₁.beforeRight etr₂).appendTrace etr₂.trace.sequence_right⟩
end Embeddings
@@ -109,21 +115,23 @@ private lemma loop_nodes_at_out :
(Graph.loop g).nodes g.loopOut = none :=
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. -/
noncomputable def EndToEndTrace.loop (etr : EndToEndTrace g ρ₁ ρ₂) :
EndToEndTrace (Graph.loop g) ρ₁ ρ₂ := by
obtain ⟨i₁, h₁, i₂, h₂, tr⟩ := etr
-- the edge in → (2 ↑ʳ i₁), reached through the second edge group
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
-- the edge (2 ↑ʳ etr.exit) → out, reached through the third edge group
have hout : (etr.exit.natAdd 2, g.loopOut) ∈ (Graph.loop g).edges := by
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 _, ?_⟩
exact Trace.single (loop_nodes_at_in ▸ EvalBasicStmtOpt.none) ++< hin >++
tr.loop ++< hout >++ Trace.single (loop_nodes_at_out ▸ EvalBasicStmtOpt.none)
exact (etr.beforeBody.appendTrace etr.trace.loop) ++< hout >++
Trace.single (loop_nodes_at_out ▸ EvalBasicStmtOpt.none)
/-- The zero-or-more times loop has an edge to return back to the top, to continue after an iteration. -/
private lemma loop_edge_out_in :
@@ -131,16 +139,23 @@ private lemma loop_edge_out_in :
refine List.mem_append_right _ ?_
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. -/
noncomputable def EndToEndTrace.loop_concat (etr₁ : EndToEndTrace (Graph.loop g) ρ₁ ρ₂)
(etr₂ : EndToEndTrace (Graph.loop g) ρ₂ ρ₃) :
EndToEndTrace (Graph.loop g) ρ₁ ρ₃ := by
obtain ⟨i₁, h₁, i₂, h₂, tr₁⟩ := etr₁
obtain ⟨j₁, k₁, j₂, k₂, tr₂⟩ := etr₂
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₂⟩
exact ⟨etr₁.entry, etr₁.entry_mem, etr₂.exit, etr₂.exit_mem,
etr₁.beforeRest etr₂ ++ etr₂.trace⟩
/-- A loop can be executed zero times. -/
noncomputable def EndToEndTrace.loop_empty {ρ : Env} : EndToEndTrace (Graph.loop g) ρ ρ := by
@@ -169,6 +184,14 @@ noncomputable def EndToEndTrace.wrap {g : Graph} {ρ₁ ρ₂ : Env}
(etr : EndToEndTrace g ρ₁ ρ₂) : EndToEndTrace (Graph.wrap g) ρ₁ ρ₂ :=
(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
possible by a language's semantics. Thus, the CFG encodes _at least_ all
semantically-possible executions. Informally, we can conclude from this

View File

@@ -126,6 +126,53 @@ theorem foldExpr_eval {vs : VariableValues ConstLattice prog} {ρ : Env} (hvs :
· 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