Skip to content

Commit 5efde9a

Browse files
committed
Computable tactic
1 parent ef5c373 commit 5efde9a

10 files changed

Lines changed: 254 additions & 9 deletions

File tree

‎RandomDo.lean‎

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,8 +11,12 @@ public import RandomDo.NumLean.PCG64
1111
public import RandomDo.NumLean.SeedSequence
1212
public import RandomDo.NumLean.Ziggurat
1313
public import RandomDo.NumLean.ZigguratSampler
14-
public import RandomDo.Tactic.Deriving
15-
public import RandomDo.Tactic.Elab
16-
public import RandomDo.Tactic.ForInStep
17-
public import RandomDo.Tactic.IsMarkov
18-
public import RandomDo.Tactic.Lemmas
14+
public import RandomDo.Tactic.Computable.Counterparts
15+
public import RandomDo.Tactic.Computable.Defs
16+
public import RandomDo.Tactic.Computable.Deriving
17+
public import RandomDo.Tactic.Computable.Example
18+
public import RandomDo.Tactic.IsMarkov.Defs
19+
public import RandomDo.Tactic.IsMarkov.Deriving
20+
public import RandomDo.Tactic.IsMarkov.Elab
21+
public import RandomDo.Tactic.IsMarkov.ForInStep
22+
public import RandomDo.Tactic.IsMarkov.Lemmas
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
public import RandomDo.Tactic.Computable.Defs
9+
public import RandomDo.NumLean.Distributions
10+
public import Mathlib.Probability.Distributions.Gaussian.Real
11+
12+
/-!
13+
# Computable counterparts of the pieces an `rdo` program is made of
14+
15+
The `@[computable]` attribute translates a program by replacing each piece it is made of by the
16+
counterpart recorded here through `@[computable_as]`. There is one entry per piece the programs of
17+
`RandomDo.Tactic.Computable.Example` are built from: the two types their values live in, and the
18+
one distribution they draw from.
19+
-/
20+
21+
public meta section
22+
23+
/-! ## Types -/
24+
25+
attribute [computable_as Float] Real
26+
attribute [computable_as Float] NNReal
27+
28+
/-! ## Distributions -/
29+
30+
/- `gaussianReal` reads its second argument as a variance and `normal` reads it as a standard
31+
deviation: the two agree on the `1` the example draws with, not in general. -/
32+
attribute [computable_as NumLean.normal] ProbabilityTheory.gaussianReal
33+
34+
end
Lines changed: 52 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,52 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
public meta import Batteries.Lean.NameMapAttribute
9+
public meta import Lean.ReservedNameAction
10+
11+
/-!
12+
# The `@[computable_as]` attribute
13+
14+
An `rdo` program is written over the Giry monad: it draws from measures on `ℝ`, which no machine
15+
samples. Turning it into a program that runs, which the `@[computable]` attribute of
16+
`RandomDo.Tactic.Computable.Deriving` does, asks for a counterpart of each piece the program is
17+
built from: `Float` for `ℝ`, `NumLean.normal` for `gaussianReal`. This file holds the attribute
18+
recording them; `RandomDo.Tactic.Computable.Counterparts` holds the counterparts themselves.
19+
20+
`@[computable_as f]` on a declaration `d` reads: `f` is what `d` becomes in a translated program.
21+
Only the pieces denoting something the program computes with need one. The scaffolding around
22+
them — numerals, arithmetic — is polymorphic, and the translation keeps it as it is, at the
23+
translated types.
24+
-/
25+
26+
public meta section
27+
28+
open Lean
29+
30+
namespace RDo.Tactic
31+
32+
/-- The counterparts recorded by `@[computable_as]`, keyed by the declaration they translate. -/
33+
initialize computableAsExt : NameMapExtension Name ←
34+
registerNameMapAttribute {
35+
name := `computable_as
36+
descr := "record the computable counterpart of this declaration"
37+
/- `@[computable_as f]` is read by `Lean.Parser.Attr.simple`, the parser an attribute that
38+
declares no syntax of its own gets: `f` is the single child of `stx[1]`. -/
39+
add := fun _ stx ↦ do
40+
let f := stx[1][0]
41+
unless f.isIdent do
42+
throwError "`computable_as` takes the name of one declaration"
43+
realizeGlobalConstNoOverload f
44+
}
45+
46+
/-- The computable counterpart of `declName`, when `@[computable_as]` recorded one. -/
47+
def computableAs? (declName : Name) : CoreM (Option Name) :=
48+
return computableAsExt.find? (← getEnv) declName
49+
50+
end RDo.Tactic
51+
52+
end
Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
public import RandomDo.Tactic.Computable.Counterparts
9+
public import RandomDo.Monad.MeasurableSpace
10+
public meta import Lean.Elab.Tactic.Basic
11+
public meta import Batteries.Tactic.Lint.Basic
12+
13+
/-!
14+
# The `@[computable]` attribute
15+
16+
Writing an `rdo` program and writing the program that samples from it are two separate steps, and
17+
the second is mechanical: every construct of the first has a counterpart in the second. This
18+
attribute walks the program and writes that counterpart down, so a program carries the sampler it
19+
denotes:
20+
21+
```
22+
@[computable]
23+
noncomputable def shifted : Measure ℝ := rdo
24+
let x ← gaussianReal 0 1
25+
return x + 1
26+
```
27+
28+
adds `shifted_computable : RandPCG IO Float`, which draws from `NumLean.normal 0 1` and adds one.
29+
30+
## How the program is read
31+
32+
The two constructs `rdo` is made of become the two of `do`: `return e` becomes `pure e`, and
33+
`let x ← p; q` becomes `p >>= q`, at `RandPCG IO`. Everything else — a distribution, a numeral, an
34+
operation on the values the program computes with — is rebuilt from the counterpart
35+
`@[computable_as]` records for its head, applied to the translation of its arguments; the
36+
instances it asks for are synthesized anew, for the translated types.
37+
-/
38+
39+
public meta section
40+
41+
open Lean Meta NumLean
42+
43+
namespace RDo.Tactic
44+
45+
/-- The monad a translated program lives in, `RandPCG IO`: the counterpart of the Giry monad an
46+
`rdo` program is written over. -/
47+
def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO]
48+
49+
/-- Translate `e`, a piece of an `rdo` program, into its computable counterpart. `σ` sends the
50+
variables the program binds to the ones the translated program binds in their place, which the
51+
change of types makes necessary. -/
52+
partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := do
53+
-- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned.
54+
if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then
55+
mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)]
56+
-- `MeasurableSpaceBind.mBind` takes eight, the last two being the program and the continuation.
57+
else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then
58+
mkAppOptM ``Bind.bind #[← computableMonad, none, none, none,
59+
← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)]
60+
else match e with
61+
| .fvar fvarId => return σ.get fvarId
62+
| .sort .. | .lit .. => return e
63+
| .mdata _ b => translate σ b
64+
| .lam .. =>
65+
lambdaBoundedTelescope e 1 fun xs body ↦ do
66+
let x := xs[0]!
67+
let t ← translate σ (← x.fvarId!.getType)
68+
withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do
69+
mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body)
70+
| _ =>
71+
let .const declName _ := e.getAppFn
72+
| throwError "`computable`: cannot translate{indentExpr e}"
73+
let counterpart := (← computableAs? declName).getD declName
74+
let mut f ← mkConstWithFreshMVarLevels counterpart
75+
for arg in e.getAppArgs do
76+
let .forallE _ t _ bi ← whnf (← inferType f)
77+
| throwError "`computable`: {f} does not take the argument{indentExpr arg}"
78+
let arg ← if bi.isInstImplicit then synthInstance t else translate σ arg
79+
unless ← isDefEq (← inferType arg) t do
80+
throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}"
81+
f := mkApp f arg
82+
return f
83+
84+
/-- Translate the `rdo` program `declName` and add the translation to the environment, under the
85+
name `declName` followed by `_computable`. -/
86+
def addComputableDecl (declName : Name) : MetaM Unit := do
87+
let info ← getConstInfo declName
88+
let some value := info.value?
89+
| throwError "`computable` can only be derived for a definition, but {declName} has no value"
90+
let value ← instantiateMVars (← translate {} value)
91+
let type ← instantiateMVars (← inferType value)
92+
let translated := declName.appendAfter "_computable"
93+
addAndCompile <| .defnDecl <| ← mkDefinitionValInferringUnsafe translated info.levelParams type
94+
value (.regular (getMaxHeight (← getEnv) value + 1))
95+
addDocStringCore translated s!"The program that samples from `{declName}`, written by the \
96+
`@[computable]` attribute."
97+
/- The name is one the attribute picks and not one the user wrote, so the underscore in it is
98+
reported for every program translated unless it is exempted here. -/
99+
setEnv (← ofExcept (Batteries.Tactic.Lint.nolintAttr.setParam (← getEnv) translated
100+
#[`defsWithUnderscore]))
101+
102+
/-- The `@[computable]` attribute. -/
103+
initialize registerBuiltinAttribute {
104+
name := `computable
105+
descr := "translate this `rdo` program into the program that samples from it"
106+
applicationTime := .afterCompilation
107+
add := fun declName _stx kind ↦ do
108+
unless kind == AttributeKind.global do
109+
throwError "`computable` must be a global attribute"
110+
(addComputableDecl declName).run'
111+
}
112+
113+
end RDo.Tactic
114+
115+
end
Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,40 @@
1+
module
2+
/- /-
3+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
4+
Released under Apache 2.0 license as described in the file LICENSE.
5+
Authors: Gaëtan Serré
6+
-/
7+
module
8+
9+
public import RandomDo.Monad.Instances
10+
public import RandomDo.Monad.Notation
11+
public import RandomDo.NumLean.Distributions
12+
public import RandomDo.Tactic.Computable.Deriving
13+
public import Mathlib
14+
meta import RandomDo.NumLean.Distributions
15+
meta import Batteries.Data.Float.Basic
16+
17+
/-!
18+
# `@[computable]` on an `rdo` program
19+
20+
`test` denotes a distribution: a Gaussian draw, shifted by one. The attribute reads it and writes
21+
the program that samples from it, drawing from `NumLean.normal` instead and shifting the draw the
22+
same way.
23+
-/
24+
25+
@[expose] public section
26+
27+
open MeasureTheory ProbabilityTheory NumLean
28+
29+
/-- Test -/
30+
@[computable]
31+
noncomputable def test (m : ℝ) : Measure ℝ := rdo
32+
let x ← gaussianReal m 1
33+
return x + 1
34+
35+
/- And it runs: seeded alike, it draws what numpy 2.3.4 draws from `default_rng(42).normal() + 1`,
36+
down to the last bit. -/
37+
run_cmd do IO.runRandPCG do
38+
let x ← (IO.runRandPCG <| test_computable 10 : IO Float)
39+
Lean.logInfo m!"`test_computable` drew {x.toStringFull}"
40+
-/
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ Authors: Rémy Degenne
55
-/
66
module
77

8-
public import RandomDo.Tactic.Elab
8+
public import RandomDo.Tactic.IsMarkov.Elab
99

1010
/-!
1111
# The `@[is_markov]` attribute
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ Authors: Gaëtan Serré
55
-/
66
module
77

8-
public import RandomDo.Tactic.Lemmas
8+
public import RandomDo.Tactic.IsMarkov.Lemmas
99
public meta import Lean.Elab.Tactic.Basic
1010

1111
/-!
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ Authors: Gaëtan Serré
55
-/
66
module
77

8-
public import RandomDo.Tactic.IsMarkov
8+
public import RandomDo.Tactic.IsMarkov.Defs
99
public import RandomDo.Monad.MeasurableSpace
1010

1111
/-!
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@ module
88
public import RandomDo.Monad.Instances
99
public import RandomDo.Monad.ForInInstances
1010
public import RandomDo.Measurable
11-
public import RandomDo.Tactic.ForInStep
11+
public import RandomDo.Tactic.IsMarkov.ForInStep
1212
public import Mathlib.MeasureTheory.Measure.ProbabilityMeasure
1313
public import Mathlib.Data.List.OfFn
1414
public import Mathlib.Probability.Distributions.Gaussian.Real

0 commit comments

Comments
 (0)