Skip to content

Commit c9192fd

Browse files
committed
Better support
1 parent 8a7666e commit c9192fd

6 files changed

Lines changed: 146 additions & 82 deletions

File tree

‎RandomDo/Tactic/Computable/Counterparts.lean‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,8 @@ attribute [computable_as Float] NNReal
2929

3030
attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal
3131

32+
/-! ## Classical functions -/
33+
3234
attribute [computable_as Float.sqrt] Real.sqrt
3335

3436
attribute [computable_as Float.log] Real.log

‎RandomDo/Tactic/Computable/Defs.lean‎

Lines changed: 2 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -11,16 +11,11 @@ public meta import Lean.ReservedNameAction
1111
/-!
1212
# The `@[computable_as]` attribute
1313
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
14+
An `rdo` program is written over the Giry monad: it draws from measures on a measurable space, which no machine samples. Turning it into a program that runs, which the `@[computable]` attribute of `RandomDo.Tactic.Computable.Deriving` does, asks for a counterpart of each piece the program is
15+
built from, e.g, `NumLean.normal'` for `gaussianReal`. This file holds the attribute
1816
recording them; `RandomDo.Tactic.Computable.Counterparts` holds the counterparts themselves.
1917
2018
`@[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.
2419
-/
2520

2621
public meta section
@@ -34,8 +29,6 @@ initialize computableAsExt : NameMapExtension Name ←
3429
registerNameMapAttribute {
3530
name := `computable_as
3631
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]`. -/
3932
add := fun _ stx ↦ do
4033
let f := stx[1][0]
4134
unless f.isIdent do

‎RandomDo/Tactic/Computable/Deriving.lean‎

Lines changed: 95 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -8,15 +8,11 @@ module
88
public import RandomDo.Tactic.Computable.Counterparts
99
public import RandomDo.Monad.MeasurableSpace
1010
public meta import Lean.Elab.Tactic.Basic
11-
public meta import Batteries.Tactic.Lint.Basic
1211

1312
/-!
1413
# The `@[computable]` attribute
1514
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:
15+
`@[computable]` reads an `rdo` program and adds the program that samples from it:
2016
2117
```
2218
@[computable]
@@ -25,19 +21,19 @@ noncomputable def shifted : Measure ℝ := rdo
2521
return x + 1
2622
```
2723
28-
adds `shifted_computable : RandPCG IO Float`, which draws from `NumLean.normal 0 1` and adds one.
24+
adds `shiftedComputable : RandPCG IO Float`, which draws from `NumLean.normal' 0 1` and adds one.
2925
30-
## How the program is read
26+
The Giry monad and its two operations become `RandPCG IO`, `pure` and `bind`. Anything else is
27+
rebuilt from the counterpart `@[computable_as]` records for its head, with its arguments translated
28+
in turn and its instances synthesized anew. A term translates into a term of the translation of its
29+
type; where the rebuilt one does not, its head is a definition nothing is known about, and its body
30+
is read in its place. `@[computable]` records the program it writes, so a program drawing from
31+
another translates into one calling that other's translation.
3132
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.
33+
Two things extend the attribute: an `@[computable_as]` entry, and an alternative of `translate` for
34+
a construct of `rdo` it has not been taught.
3735
38-
Setting `set_option trace.computable true` prints the tree of pieces the attribute walked through,
39-
each with what it became, the instances synthesized along the way, and the declaration written at
40-
the end.
36+
`set_option trace.computable true` prints what each piece became.
4137
-/
4238

4339
public meta section
@@ -48,77 +44,104 @@ namespace RDo.Tactic
4844

4945
initialize registerTraceClass `computable
5046

51-
/-- The monad a translated program lives in, `RandPCG IO`: the counterpart of the Giry monad an
52-
`rdo` program is written over. -/
47+
/-- The monad the translated programs live in. -/
5348
def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO]
5449

55-
/-- Translate `e`, a piece of an `rdo` program, into its computable counterpart. `σ` sends the
56-
variables the program binds to the ones the translated program binds in their place, which the
57-
change of types makes necessary. -/
50+
mutual
51+
52+
/-- Translate a piece of an `rdo` program. `σ` maps the variables the program binds to the ones the
53+
translated program binds in their place. -/
5854
partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr :=
59-
withTraceNode `computable
60-
(fun
61-
| .ok e' => return m!"{e} ↦ {e'}"
62-
| .error _ => return m!"{e}: not translated") do
63-
-- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned.
64-
if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then
65-
mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)]
66-
-- `MeasurableSpaceBind.mBind` takes eight, the last two the program and the continuation.
67-
else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then
55+
withTraceNode `computable (fun
56+
| .ok e' => return m!"{e} ↦ {e'}"
57+
| .error _ => return m!"{e}: not translated") do
58+
match_expr e with
59+
| MeasurableSpacePure.mPure _ _ _ _ a =>
60+
mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ a]
61+
| MeasurableSpaceBind.mBind _ _ _ _ _ _ p k =>
6862
mkAppOptM ``Bind.bind #[← computableMonad, none, none, none,
69-
← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)]
70-
else match e with
71-
| .fvar fvarId => return σ.get fvarId
72-
| .sort .. | .lit .. => return e
73-
| .mdata _ b => translate σ b
74-
| .lam .. =>
75-
lambdaBoundedTelescope e 1 fun xs body ↦ do
76-
let x := xs[0]!
77-
let t ← translate σ (← x.fvarId!.getType)
78-
withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do
79-
mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body)
80-
| _ =>
81-
let .const declName _ := e.getAppFn
82-
| throwError "`computable`: cannot translate{indentExpr e}"
83-
let counterpart := (← computableAs? declName).getD declName
84-
let mut f ← mkConstWithFreshMVarLevels counterpart
85-
for arg in e.getAppArgs do
86-
let .forallE _ t _ bi ← whnf (← inferType f)
87-
| throwError "`computable`: {f} does not take the argument{indentExpr arg}"
88-
/- An instance is the one argument that is not translated but synthesized anew, so the
89-
trace above it says nothing: it is reported here instead. -/
90-
let arg ←
91-
if bi.isInstImplicit then do
92-
let inst ← synthInstance t
93-
trace[computable] "instance: {inst}"
94-
pure inst
95-
else
96-
translate σ arg
97-
unless ← isDefEq (← inferType arg) t do
98-
throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}"
99-
f := mkApp f arg
100-
return f
63+
← translate σ p, ← translate σ k]
64+
| MeasureTheory.Measure α _ => return mkApp (← computableMonad) (← translate σ α)
65+
| _ => match e with
66+
| .fvar x => return σ.get x
67+
| .sort .. | .lit .. => return e
68+
| .mdata _ b => translate σ b
69+
| .lam .. => lambdaBoundedTelescope e 1 fun xs body ↦ do
70+
let x := xs[0]!.fvarId!
71+
withLocalDeclD (← x.getUserName) (← translate σ (← x.getType)) fun y ↦ do
72+
mkLambdaFVars #[y] (← translate (σ.insert x y) body)
73+
| _ => translateApp σ e
74+
75+
/-- Rebuild an application from the counterpart of its head; where nothing known about that head
76+
fits, look through it and read its body in its place. -/
77+
partial def translateApp (σ : FVarSubst) (e : Expr) : MetaM Expr := do
78+
if e.getAppFn.isLambda then return ← translate σ e.headBeta
79+
let .const declName _ := e.getAppFn | throwError "`computable`: cannot translate{indentExpr e}"
80+
let counterpart? ← computableAs? declName
81+
let head := counterpart?.getD declName
82+
/- We save the state of metavariables, so that a failed rebuild does not leave them in a
83+
half-built state. -/
84+
let s ← saveState
85+
try
86+
/- A term translates into a term of the translation of its type. A head nothing is known about
87+
rebuilds into itself, of the type it had, and might fail. We check that the rebuilt term has
88+
the translation of the type, and if not, we look through it. -/
89+
let f ← rebuild σ head e
90+
-- Useless for a type, whose type is a sort either way.
91+
if (← inferType f).isSort then return f
92+
let expected ← translate σ (← inferType e)
93+
unless ← isDefEq (← inferType f) expected do
94+
throwError "`computable`: {f} is of type{indentExpr (← inferType f)}\n\
95+
where the translation asks for{indentExpr expected}"
96+
return f
97+
catch ex =>
98+
restoreState s
99+
/- The rebuild failed, we try to look through the head, and read its body in its place. -/
100+
if counterpart?.isNone then
101+
if let some e' ← unfoldDefinition? e then
102+
trace[computable] "nothing known about {declName}, looking through it"
103+
return ← translate σ e'
104+
throw ex
105+
106+
/-- Apply `head` to the arguments of `e`, translated in turn, and check that the result has the
107+
translation of the type of `e`. -/
108+
partial def rebuild (σ : FVarSubst) (head : Name) (e : Expr) : MetaM Expr := do
109+
let mut f ← mkConstWithFreshMVarLevels head
110+
for arg in e.getAppArgs do
111+
let .forallE _ t _ bi ← whnf (← inferType f)
112+
| throwError "`computable`: {f} does not take the argument{indentExpr arg}"
113+
let arg ←
114+
if bi.isInstImplicit then do
115+
-- An instance is not translated: it is asked for anew, at the translated types.
116+
let inst ← synthInstance t
117+
trace[computable] "instance: {inst}"
118+
pure inst
119+
else
120+
translate σ arg
121+
unless ← isDefEq (← inferType arg) t do
122+
throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}"
123+
f := mkApp f arg
124+
return f
125+
126+
end
101127

102128
/-- Translate the `rdo` program `declName` and add the translation to the environment, under the
103-
name `declName` followed by `_computable`. -/
129+
name `declName` followed by `Computable`. -/
104130
def addComputableDecl (declName : Name) : MetaM Unit := do
105131
let info ← getConstInfo declName
106132
let some value := info.value?
107133
| throwError "`computable` can only be derived for a definition, but {declName} has no value"
108134
let value ← instantiateMVars (← translate {} value)
109135
let type ← instantiateMVars (← inferType value)
110-
let translated := declName.appendAfter "_computable"
136+
let translated := declName.appendAfter "Computable"
111137
addAndCompile <| .defnDecl <| ← mkDefinitionValInferringUnsafe translated info.levelParams type
112138
value (.regular (getMaxHeight (← getEnv) value + 1))
113-
trace[computable] "wrote {translated} :{indentExpr type}"
114-
addDocStringCore translated s!"The program that samples from `{declName}`, written by the \
115-
`@[computable]` attribute."
116-
/- The name is one the attribute picks and not one the user wrote, so the underscore in it is
117-
reported for every program translated unless it is exempted here. -/
118-
setEnv (← ofExcept (Batteries.Tactic.Lint.nolintAttr.setParam (← getEnv) translated
119-
#[`defsWithUnderscore]))
120-
121-
/-- The `@[computable]` attribute. -/
139+
addDocStringCore translated s!"The computable program that samples from `{declName}` \
140+
(automatically generated by the `@[computable]` attribute)."
141+
computableAsExt.add declName translated
142+
trace[computable] "wrote {translated}:{indentExpr type}"
143+
144+
@[inherit_doc addComputableDecl]
122145
initialize registerBuiltinAttribute {
123146
name := `computable
124147
descr := "translate this `rdo` program into the program that samples from it"

‎RandomDo/Tactic/IsMarkov/Deriving.lean‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -63,7 +63,7 @@ def addIsMarkovInstance (declName : Name) : TermElabM Unit := do
6363
value := ← instantiateMVars proof })
6464
Meta.addInstance instName .global 1000
6565

66-
/-- The `@[is_markov]` attribute. -/
66+
@[inherit_doc isMarkovStatement]
6767
initialize registerBuiltinAttribute {
6868
name := `is_markov
6969
descr := "prove that this `rdo` program is a Markov kernel, and register it as an instance"

‎Test.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ module -- shake: keep-all --deprecated_module: ignore
22

33
public import Test.Bind
44
public import Test.Common
5+
public import Test.Computable
56
public import Test.Control
67
public import Test.Gaps
78
public import Test.Instances

‎Test/Computable.lean‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
module
2+
3+
public import Test.IsMarkov
4+
import Batteries.Data.Float.Basic
5+
/- A `run_cmd` runs at elaboration time, so what it calls has to be imported as `meta` too: the
6+
sampler it draws with, and `Float.toStringFull` it prints with. -/
7+
meta import RandomDo.NumLean.Distributions
8+
meta import Batteries.Data.Float.Basic
9+
10+
set_option linter.style.header false
11+
12+
set_option trace.computable true
13+
14+
namespace Test.Computable
15+
16+
open Test.IsMarkov NumLean Lean.Elab.Command
17+
18+
def logComputable (prog : RandPCG IO Float) : CommandElabM Unit := do
19+
let x ← (IO.runRandPCG prog : IO Float)
20+
let y ← (IO.runRandPCGWith 42 prog : IO Float)
21+
Lean.logInfo m!"x = {x.toStringFull}"
22+
Lean.logInfo m!"y (seed 42) = {y.toStringFull}"
23+
24+
--attribute [computable] sumTwo
25+
26+
--run_cmd do logComputable sumTwoComputable
27+
28+
@[computable]
29+
noncomputable
30+
def test : MeasureTheory.Measure ℝ := rdo
31+
let y ← sumTwo
32+
let x ← ProbabilityTheory.gaussianReal 0 1
33+
return x + y
34+
35+
attribute [computable] centred
36+
37+
run_cmd do logComputable (centredComputable 20)
38+
39+
attribute [computable] branchOn
40+
41+
run_cmd do logComputable (branchOnComputable 20)
42+
43+
run_cmd do logComputable (branchOnComputable (-1))
44+
45+
end Test.Computable

0 commit comments

Comments
 (0)