diff --git a/lean/Spa/Transformation/Licm.lean b/lean/Spa/Transformation/Licm.lean index 51f8654..be6b7d7 100644 --- a/lean/Spa/Transformation/Licm.lean +++ b/lean/Spa/Transformation/Licm.lean @@ -30,6 +30,9 @@ assignment to lead the body, its destination to be absent from the guard, and 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. -/ diff --git a/lean/Spa/Transformation/Licm/Correctness.lean b/lean/Spa/Transformation/Licm/Correctness.lean index 7ea7f4b..6156996 100644 --- a/lean/Spa/Transformation/Licm/Correctness.lean +++ b/lean/Spa/Transformation/Licm/Correctness.lean @@ -76,5 +76,184 @@ lemma head_rhs_agrees {prog : Program} {cond : Expr} {x : String} {rhs : Expr} { (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