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

132 lines
5.4 KiB
Lean4
Raw Normal View History

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