Files
agda-spa/lean/Spa/Transformation/Constant.lean

132 lines
5.4 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.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₂)
end ConstantTransform
end Spa