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