Files
agda-spa/lean/Spa/Transformation/Licm/Correctness.lean

260 lines
14 KiB
Lean4
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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