260 lines
14 KiB
Lean4
260 lines
14 KiB
Lean4
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
|