Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions LeanMachineLearning.lean
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,11 @@ public import LeanMachineLearning.Online.Bandit.BayesRegret
public import LeanMachineLearning.Online.Bandit.Regret
public import LeanMachineLearning.Online.Bandit.RewardByCountMeasure
public import LeanMachineLearning.Online.Bandit.SumRewards
public import LeanMachineLearning.RDo.Examples
public import LeanMachineLearning.RDo.ForInInstances
public import LeanMachineLearning.RDo.MeasurableSpaceMonad
public import LeanMachineLearning.RDo.MonadInstances
public import LeanMachineLearning.RDo.RDo
public import LeanMachineLearning.SequentialLearning.Algorithm
public import LeanMachineLearning.SequentialLearning.AlgorithmDensity
public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes
Expand Down
57 changes: 57 additions & 0 deletions LeanMachineLearning/RDo/Examples.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
module

public import LeanMachineLearning.RDo.MonadInstances
public import LeanMachineLearning.RDo.ForInInstances
public import Mathlib.Probability.Distributions.Bernoulli
public import Mathlib.Algebra.Ring.BooleanRing

set_option linter.style.header false

@[expose] public section

open MeasureTheory ProbabilityTheory Measure

/- # Nonpolymorphic examples -/

universe u v

noncomputable def measureSample : Measure Bool := rdo
let x ← bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩
let y ← bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩
return x + y

def pseudoSample : PseudoRandomM Bool := rdo
let x ← Random.randBool
let y ← Random.randBool
return x + y

/- # Polymorphic examples -/

variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m]

class HasBit (m : (α : Type) → MeasurableSpace α → Type v) where
bit : m Bool (by infer_instance)

noncomputable instance : HasBit Measure where
bit := bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩

instance : HasBit PseudoRandomM where
bit := Random.randBool

def indepAnd [HasBit m] : m Bool := rdo
let x ← HasBit.bit
let y ← HasBit.bit
return x && y

noncomputable def indepAndMeasure : Measure Bool := indepAnd (m := Measure)

def indepAndGen : PseudoRandomM Bool := indepAnd (m := PseudoRandomM)

variable {α : Type*} [MeasurableSpace α]

def sampleBitsArray [HasBit m] (n : ℕ) : m (Array Bool) := rdo
let mut xs : Array Bool := #[]
for _ in List.range n rdo
let b ← HasBit.bit (m := m)
xs := xs.push b
return xs
110 changes: 110 additions & 0 deletions LeanMachineLearning/RDo/ForInInstances.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
/-
Copyright (c) 2026 David Ledvinka. All rights reserved.
Released under Apache 2.0 license as described in the file LICENSE.
Authors: David Ledvinka
-/
module

public import LeanMachineLearning.RDo.RDo

/-!
# Instances for `MeasurableSpaceForIn`

**TODO**

-/

@[expose] public section

section MeasurableSpace

variable {α : Type*} [MeasurableSpace α]

instance : MeasurableSpace (List α) :=
MeasurableSpace.comap List.equivSigmaTuple inferInstance

instance : MeasurableSpace (Array α) :=
MeasurableSpace.comap Array.toList inferInstance

end MeasurableSpace

universe u v

open MeasurableSpacePure

variable {α : Type u} [mα : MeasurableSpace α] {m : (α : Type u) → [MeasurableSpace α] → Type v}

section Array

/-- Compiler implementation for `forIn` -/
@[inline] unsafe def Array.measurableSpaceForIn'Unsafe [MeasurableSpaceMonad m]
{β : Type u} [mβ : MeasurableSpace β]
(as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
let sz := as.usize
let rec @[specialize] loop (i : USize) (b : β) : m β := rdo
if i < sz then
let a := as.uget i lcProof
match (← f a lcProof b) with
| ForInStep.done b => mPure b
| ForInStep.yield b => loop (i+1) b
else
mPure b
loop 0 b

/-- Reference implementation for `forIn'` -/
@[implemented_by Array.measurableSpaceForIn'Unsafe]
protected def Array.measurableSpaceForIn' [MeasurableSpaceMonad m]
{β : Type u} [mβ : MeasurableSpace β]
(as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
let rec loop (i : Nat) (h : i ≤ as.size) (b : β) : m β := rdo
match i, h with
| 0, _ => mPure b
| i+1, h =>
have h' : i < as.size := Nat.lt_of_lt_of_le (Nat.lt_succ_self i) h
have : as.size - 1 < as.size := Nat.sub_lt (Nat.zero_lt_of_lt h') (by decide)
have : as.size - 1 - i < as.size := Nat.lt_of_le_of_lt (Nat.sub_le (as.size - 1) i) this
match (← f as[as.size - 1 - i] (getElem_mem this) b) with
| ForInStep.done b => mPure b
| ForInStep.yield b => loop i (Nat.le_of_lt h') b
loop as.size (Nat.le_refl _) b

instance [MeasurableSpaceMonad m] : MeasurableSpaceForIn' m (Array α) α inferInstance where
forIn' := Array.measurableSpaceForIn'

instance {n : ℕ} [MeasurableSpaceMonad m] :
MeasurableSpaceForIn' m (Vector α n) α inferInstance where
forIn' xs b f := Array.measurableSpaceForIn' xs.toArray b (fun a h b => f a (by simpa using h) b)

end Array

section List

variable {α β : Type*} [MeasurableSpace α] [MeasurableSpace β] [Ring α]

/-- Implimentation for `forIn'` -/
@[inline]
protected def List.measurableSpaceForIn' [MeasurableSpaceMonad m]
{β : Type u} [mβ : MeasurableSpace β] (as : @& List α) (init : β)
(f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
let rec @[specialize]
loop : (as' : @& List α) → (b : β) → Exists (fun bs => bs ++ as' = as) → m β
| [], b, _ => mPure b
| a::as', b, h => rdo
have : a ∈ as := by
clear f
have ⟨bs, h⟩ := h
subst h
exact mem_append_right _ (Mem.head ..)
match (← f a this b) with
| ForInStep.done b => mPure b
| ForInStep.yield b =>
have : Exists (fun bs => bs ++ as' = as) :=
have ⟨bs, h⟩ := h
⟨bs ++ [a], by rw [← h, append_cons (bs := as')]⟩
loop as' b this
loop as init ⟨[], rfl⟩

instance [MeasurableSpaceMonad m] : MeasurableSpaceForIn' m (List α) α inferInstance where
forIn' := List.measurableSpaceForIn'

end List
Loading
Loading