179 lines
7.7 KiB
Lean4
179 lines
7.7 KiB
Lean4
import Spa.Language.Base
|
||
import Spa.Language.Program
|
||
import Spa.Lattice.FiniteMap
|
||
import Spa.Analysis.Constant
|
||
import Spa.Analysis.Forward
|
||
|
||
/-!
|
||
# Constant folding
|
||
|
||
Rewrites each assignment's right-hand side to a literal wherever the constant
|
||
analysis (`Spa/Analysis/Constant.lean`) pins its value down.
|
||
|
||
The traversal recurses over the plain `Stmt`, threading a `GGraph.Embed` of the
|
||
current subtree's CFG into the program's (`Program.rootEmbed`, then one
|
||
`Embed.trans` per descent). At an assignment, `Embed.singletonIndex` gives its
|
||
CFG state, and the facts to fold with are `Forward.joinForKey` at that state —
|
||
the join over predecessors, i.e. the values *entering* the node, which is what
|
||
the right-hand side reads. (`Forward.variablesAt` would be the values *leaving*
|
||
it, which already include this assignment's own effect.)
|
||
-/
|
||
|
||
namespace Spa
|
||
|
||
namespace ConstantTransform
|
||
|
||
open GGraph Forward
|
||
|
||
variable (prog : Program)
|
||
|
||
abbrev result := Forward.result ConstLattice prog
|
||
|
||
/-- Replace an expression by a literal when the analysis pins its value down,
|
||
recursing into its subexpressions otherwise. Whole subexpressions are tried
|
||
first, so `(x + 1) - x` folds outright when `x` is known, rather than only in its
|
||
leaves. -/
|
||
def foldExpr (vs : VariableValues ConstLattice prog) : Expr → Expr
|
||
| .num n => .num n
|
||
| .var k =>
|
||
match ConstAnalysis.eval prog (.var k) vs with
|
||
| .mk z => .num z
|
||
| _ => .var k
|
||
| .add a b =>
|
||
match ConstAnalysis.eval prog (.add a b) vs with
|
||
| .mk z => .num z
|
||
| _ => .add (foldExpr vs a) (foldExpr vs b)
|
||
| .sub a b =>
|
||
match ConstAnalysis.eval prog (.sub a b) vs with
|
||
| .mk z => .num z
|
||
| _ => .sub (foldExpr vs a) (foldExpr vs b)
|
||
|
||
/-- Constant-fold every assignment in a statement.
|
||
|
||
`sv` is the analysis result, taken as a parameter rather than read from `result`
|
||
at each node: it is a fixpoint computation, so recomputing it per assignment
|
||
would make folding quadratic in the analysis.
|
||
|
||
Guards of `ifElse`/`whileLoop` are deliberately left alone. `Stmt.cfg` gives a
|
||
conditional's guard no node at all (`ifElse` overlays the two branches), so there
|
||
is no state whose entry facts describe where it is evaluated. A `whileLoop`'s
|
||
guard does have a candidate — the loop header `GGraph.loopIn` — but tying the
|
||
guard's evaluation environment to that node needs a lemma that does not exist
|
||
yet, so folding it here would be an unproven soundness claim. -/
|
||
def foldStmt (sv : StateVariables ConstLattice prog) :
|
||
(s : Stmt) → Embed s.cfg prog.cfg → Stmt
|
||
| .basic .noop, _ => .basic .noop
|
||
| .basic (.assign k v), e =>
|
||
.basic (.assign k (foldExpr prog (joinForKey e.singletonIndex sv) v))
|
||
| .andThen s₁ s₂, e =>
|
||
.andThen (foldStmt sv s₁ ((Embed.sequenceLeft s₁.cfg s₂.cfg).trans e))
|
||
(foldStmt sv s₂ ((Embed.sequenceRight s₁.cfg s₂.cfg).trans e))
|
||
| .ifElse cond s₁ s₂, e =>
|
||
.ifElse cond (foldStmt sv s₁ ((Embed.overlayLeft s₁.cfg s₂.cfg).trans e))
|
||
(foldStmt sv s₂ ((Embed.overlayRight s₁.cfg s₂.cfg).trans e))
|
||
| .whileLoop cond body, e =>
|
||
.whileLoop cond (foldStmt sv body ((Embed.loop body.cfg).trans e))
|
||
|
||
/-- Constant-fold a whole program, running the analysis once. -/
|
||
def foldProgram : Stmt := foldStmt prog (result prog) prog.rootStmt prog.rootEmbed
|
||
|
||
/-! ## Correctness
|
||
|
||
Folding preserves meaning *provided the facts folded with actually hold of the
|
||
environment folded in*. That proviso is the whole content: `foldExpr` is sound
|
||
against any `vs` that over-approximates `ρ`, and it is the analysis engine's job
|
||
(`Forward.analyze_correct_at`) to supply such a `vs` at each program point. -/
|
||
|
||
variable {prog}
|
||
|
||
/-- If the analysis pins an expression to a constant and its facts hold of `ρ`,
|
||
then the expression really does evaluate to that constant. This is
|
||
`ValidExprEvaluator` specialised to the `.mk` case, where `interpConst` says
|
||
exactly `v = .int z`. -/
|
||
lemma eq_int_of_eval_mk {vs : VariableValues ConstLattice prog} {ρ : Env}
|
||
{e : Expr} {v : Value} {z : ℤ}
|
||
(hev : EvalExpr ρ e v) (hvs : ⟦vs⟧ ρ) (hz : ConstAnalysis.eval prog e vs = .mk z) :
|
||
v = .int z := by
|
||
have h := ValidExprEvaluator.valid (L := ConstLattice) (prog := prog) hev hvs
|
||
rw [show ExprEvaluator.eval e vs = ConstAnalysis.eval prog e vs from rfl, hz] at h
|
||
exact h
|
||
|
||
/-- **Expression folding is meaning-preserving.** Whenever `vs` over-approximates
|
||
`ρ`, the folded expression evaluates in `ρ` to whatever the original did. -/
|
||
theorem foldExpr_eval {vs : VariableValues ConstLattice prog} {ρ : Env} (hvs : ⟦vs⟧ ρ) :
|
||
∀ {e : Expr} {v : Value}, EvalExpr ρ e v → EvalExpr ρ (foldExpr prog vs e) v := by
|
||
intro e
|
||
induction e with
|
||
| num n => intro v hev; simpa [foldExpr] using hev
|
||
| var k =>
|
||
intro v hev
|
||
simp only [foldExpr]
|
||
split
|
||
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
|
||
· exact hev
|
||
| add a b iha ihb =>
|
||
intro v hev
|
||
simp only [foldExpr]
|
||
split
|
||
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
|
||
· cases hev with
|
||
| add _ _ z₁ z₂ h₁ h₂ => exact EvalExpr.add ρ _ _ z₁ z₂ (iha h₁) (ihb h₂)
|
||
| sub a b iha ihb =>
|
||
intro v hev
|
||
simp only [foldExpr]
|
||
split
|
||
· case h_1 z hz => rw [eq_int_of_eval_mk hev hvs hz]; exact EvalExpr.num ρ z
|
||
· 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
|