Skip to content
Merged
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
12 changes: 8 additions & 4 deletions LeanMachineLearning/Online/Bandit/ArrayProbSpace.lean
Original file line number Diff line number Diff line change
Expand Up @@ -477,13 +477,15 @@ lemma truePast_eq_of_pullCount_eq (alg : Algorithm 𝓐 R)
truePast alg a n ω = (ω.1, fun i b ↦ if b = a then if m ≠ 0 then
ω.2 (min i (m - 1)) a else Nonempty.some inferInstance else ω.2 i b) := by
simp [truePast, h_pc]
grind

lemma truePast_eq_of_pullCount_eq_of_ne_zero (alg : Algorithm 𝓐 R)
(a : 𝓐) (n m : ℕ) (ω : probSpace 𝓐 R)
(h_pc : pullCount (action alg) a (n + 1) ω = m) (hm : m ≠ 0) :
truePast alg a n ω = (ω.1, fun i b ↦ if b = a then
ω.2 (min i (m - 1)) a else ω.2 i b) := by
simp [truePast, h_pc, hm]
grind

lemma measurable_hist_truePast [Countable 𝓐] (alg : Algorithm 𝓐 R)
(a : 𝓐) (n : ℕ) :
Expand Down Expand Up @@ -569,12 +571,14 @@ lemma measurable_pullCount_action_add_one_hist (alg : Algorithm 𝓐 R) (n : ℕ

end MeasurabilityAdvanced

set_option backward.isDefEq.respectTransparency false in
omit [Nonempty 𝓐] [StandardBorelSpace 𝓐] [DecidableEq 𝓐] in
lemma map_snd_apply_arrayMeasure {ν : Kernel 𝓐 R} [IsMarkovKernel ν] (n : ℕ) (a : 𝓐) :
(arrayMeasure ν).map (fun ω ↦ ω.2 n a) = ν a := by
calc (arrayMeasure ν).map (fun ω ↦ ω.2 n a)
_ = (arrayMeasure ν).snd.map (fun ω ↦ ω n a) := by
rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]
unfold Measure.snd
rw [Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = ν a := by
rw [arrayMeasure, Measure.snd_prod, streamMeasure]
Expand Down Expand Up @@ -653,6 +657,7 @@ lemma indepFun_fst_add_one_hist [Countable 𝓐] (alg : Algorithm 𝓐 R)
(indepFun_fst_add_one_aux ν n).of_measurable_right (measurable_hist_comap alg n)

-- proved by Claude
set_option backward.isDefEq.respectTransparency false in
omit [Nonempty 𝓐] [StandardBorelSpace 𝓐] [StandardBorelSpace R] in
lemma indepFun_snd_apply_aux (ν : Kernel 𝓐 R) [IsMarkovKernel ν] (a : 𝓐) (m : ℕ) :
(fun ω ↦ ω.2 m a) ⟂ᵢ[arrayMeasure ν]
Expand Down Expand Up @@ -912,9 +917,7 @@ lemma indepFun_snd_hist_cond [Countable 𝓐] (alg : Algorithm 𝓐 R)
fun ω ↦ (ω.1, fun k b ↦ if b = a then if m ≠ 0 then ω.2 (min k (m - 1)) b
else Nonempty.some inferInstance else ω.2 k b) by
convert this using 1
· rfl
· rfl
· rfl
congr!
congr with ω
simp only [Set.mem_preimage, Set.mem_singleton_iff, Prod.mk.injEq, Set.indicator_apply,
Set.mem_setOf_eq, ite_eq_left_iff, not_and, zero_ne_one, imp_false,
Expand All @@ -938,6 +941,7 @@ section Laws

variable [Countable 𝓐]

set_option backward.isDefEq.respectTransparency false in
lemma hasLaw_action_zero (alg : Algorithm 𝓐 R) (ν : Kernel 𝓐 R) [IsMarkovKernel ν] :
HasLaw (action alg 0) alg.p0 (arrayMeasure ν) where
map_eq := by
Expand Down
2 changes: 2 additions & 0 deletions LeanMachineLearning/Online/Bandit/SumRewards.lean
Original file line number Diff line number Diff line change
Expand Up @@ -143,12 +143,14 @@ lemma sumRewards_eq_comp :
(fun p ↦ ∑ i ∈ range n, if (p i).1 = a then (p i).2 else 0) ∘ (trajectory A R) := by
ext
simp [sumRewards, trajectory]
grind

lemma pullCount_eq_comp :
pullCount A a n =
(fun p ↦ ∑ i ∈ range n, if (p i).1 = a then 1 else 0) ∘ (trajectory A R) := by
ext
simp [pullCount, trajectory]
rfl

-- todo: write those lemmas with IdentDistrib instead of equality of maps
lemma _root_.Learning.IsAlgEnvSeq.law_sumRewards_unique [MeasurableSingletonClass 𝓐]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,10 @@ lemma sum_mod_range_mul {K : ℕ} (hK : 0 < K) (m : ℕ) (a : Fin K) :
_ = n + (∑ s ∈ range K, if ⟨(s + K * n) % K, Nat.mod_lt _ hK⟩ = a then 1 else 0) := by
congr 1
let e : ℕ ↪ ℕ := ⟨fun i : ℕ ↦ i + K * n, fun i j hij ↦ by grind⟩
have : Finset.map e (range K) = Ico (K * n) (K * n + K) := by
have he i : e i = i + K * n := rfl
have : (range K).map e = Ico (K * n) (K * n + K) := by
ext x
simp only [mem_map, mem_range, Function.Embedding.coeFn_mk, mem_Ico, e]
simp only [mem_map, mem_range, mem_Ico, e]
refine ⟨fun h ↦ by grind, fun h ↦ ?_⟩
use x - K * n
grind
Expand Down
3 changes: 2 additions & 1 deletion LeanMachineLearning/SequentialLearning/FiniteActions.lean
Original file line number Diff line number Diff line change
Expand Up @@ -696,7 +696,8 @@ lemma rewardByCount_of_stepsUntil_ne_top (h : stepsUntil A a m ω.1 ≠ ⊤) :

lemma rewardByCount_eq_stoppedValue (h : stepsUntil A a m ω.1 ≠ ⊤) :
rewardByCount A R' a m ω = stoppedValue R' (stepsUntil A a m) ω.1 := by
rw [rewardByCount_of_stepsUntil_ne_top h, stoppedValue]
unfold stoppedValue
rw [rewardByCount_of_stepsUntil_ne_top h]
lift stepsUntil A a m ω.1 to ℕ using h with n
simp

Expand Down
28 changes: 14 additions & 14 deletions lake-manifest.json
Original file line number Diff line number Diff line change
Expand Up @@ -5,17 +5,17 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "11d11a11a667a8fa8ea19d9456fe059f683e308f",
"rev": "79d0395a1825a6264ad5d269e35e60537518955e",
"name": "mathlib",
"manifestFile": "lake-manifest.json",
"inputRev": "11d11a11a667a8fa8ea19d9456fe059f683e308f",
"inputRev": "79d0395a1825a6264ad5d269e35e60537518955e",
"inherited": false,
"configFile": "lakefile.lean"},
{"url": "https://github.com/leanprover/verso",
"type": "git",
"subDir": null,
"scope": "",
"rev": "c741b6f1fa5328dd0a85fab08c87de119b7fe4fe",
"rev": "6af30619664960f9c816b27157dff3ceb1500dcf",
"name": "verso",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -25,7 +25,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "e12c1910fe855cbfc38803cd4e55543906d5fa62",
"rev": "b1c4a69a7e247ab7df20460212001673d74f08c0",
"name": "plausible",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -45,7 +45,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "7e9612bf0b9ee66db3cb5b9988a35afc706f5a12",
"rev": "18a90119a5d316358fde6c86e0ca24e59212e32c",
"name": "importGraph",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -55,7 +55,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "d662197a9ca6f411c5738c45ec0192c786462f5d",
"rev": "b1436dc749e722c9920036b52cdc43b3451d0b69",
"name": "proofwidgets",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -65,7 +65,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "a7dbf0c63b694e47f425f3dcddbc0e178bb432d3",
"rev": "57d3325be72a842920813bcb40f96a6f7393c185",
"name": "aesop",
"manifestFile": "lake-manifest.json",
"inputRev": "master",
Expand All @@ -75,7 +75,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "38d591e778f100aec9762bb582f9c7f55f50e9dc",
"rev": "ee41917ae11d38479fb8fb24745f7ca4bf0a784d",
"name": "Qq",
"manifestFile": "lake-manifest.json",
"inputRev": "master",
Expand All @@ -85,7 +85,7 @@
"type": "git",
"subDir": null,
"scope": "leanprover-community",
"rev": "023ce7d62a0531e22a5331e20b587817a80d49ff",
"rev": "31a49105f960721073a9adfc82b261f5d0f2ce1e",
"name": "batteries",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -95,7 +95,7 @@
"type": "git",
"subDir": null,
"scope": "",
"rev": "ae95e7e7d01c072421732d0b84cf63ff903f4f0e",
"rev": "56958b3901ca108830de34fbce6cecd4b5757c1f",
"name": "illuminate",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -105,7 +105,7 @@
"type": "git",
"subDir": null,
"scope": "",
"rev": "6a3fb240133bcb7e1a066fdc784b3fdc304e3fc5",
"rev": "31907cc18f48a95384f99cee5582c00fb39e0f67",
"name": "MD4Lean",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -115,7 +115,7 @@
"type": "git",
"subDir": null,
"scope": "",
"rev": "0bd508e8362f56d4a05cbf63614d4c97db954041",
"rev": "0076a9e8a3670d83c54c93414b2b26d3a8aba08d",
"name": "subverso",
"manifestFile": "lake-manifest.json",
"inputRev": "main",
Expand All @@ -125,10 +125,10 @@
"type": "git",
"subDir": null,
"scope": "leanprover",
"rev": "88679d088c9720c27ebdf2ba4dafe17341747f94",
"rev": "da07ca808b6718cb2aed14dba154e5a08b8f8ecf",
"name": "Cli",
"manifestFile": "lake-manifest.json",
"inputRev": "v4.32.0",
"inputRev": "v4.33.0-rc1",
"inherited": true,
"configFile": "lakefile.toml"}],
"name": "LeanMachineLearning",
Expand Down
2 changes: 1 addition & 1 deletion lakefile.toml
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ rev = "main"
[[require]]
name = "mathlib"
scope = "leanprover-community"
rev = "11d11a11a667a8fa8ea19d9456fe059f683e308f"
rev = "79d0395a1825a6264ad5d269e35e60537518955e"

[[lean_lib]]
name = "LeanMachineLearning"
Expand Down
2 changes: 1 addition & 1 deletion lean-toolchain
Original file line number Diff line number Diff line change
@@ -1 +1 @@
leanprover/lean4:v4.32.0
leanprover/lean4:v4.33.0-rc1