|
| 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 |
0 commit comments