Skip to content
Closed
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
127 changes: 127 additions & 0 deletions Cslib/Probability/PMF.lean
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,9 @@ the Mathlib module instead.
- `Cslib.Probability.PMF.bind_pair_tsum_fst`: marginalizing over the first component
- `Cslib.Probability.PMF.uniformOfFintype_map_equiv`:
a uniform distribution is invariant under equivalence
- `Cslib.Probability.PMF.uniformOfFintype_prod`: a uniform pair is two independent uniform samples
- `Cslib.Probability.PMF.toOuterMeasure_bind_failure_le`: errors compose through a
state-dependent continuation without independence assumptions
- `Cslib.Probability.PMF.posteriorDist`: the posterior as a `PMF`
- `Cslib.Probability.PMF.posteriorDist_eq_prior_of_outputIndist`:
if the output distribution does not depend on the input, conditioning does
Expand All @@ -51,6 +54,32 @@ theorem summable_toReal (p : PMF α) : Summable (fun a => (p a).toReal) :=
@[simp] theorem tsum_toReal (p : PMF α) : ∑' a, (p a).toReal = 1 := by
rw [← ENNReal.tsum_toReal_eq p.apply_ne_top, p.tsum_coe, ENNReal.toReal_one]

/-- An injective representation of outcomes also embeds their complete probability laws. -/
theorem map_injective {f : α → β} (hf : Function.Injective f) :
Function.Injective (PMF.map f) := by
classical
intro p q h
ext a
simpa [hf.eq_iff] using DFunLike.congr_fun h (f a)

/-- Relabeling a distribution by an equivalence preserves each corresponding point mass. -/
theorem map_equiv_apply (p : PMF α) (e : α ≃ β) (b : β) :
p.map e b = p (e.symm b) := by
classical
simp [← e.symm_apply_eq]

/-- A deterministic map preserves at least the mass of each individual preimage. -/
theorem le_map_apply (p : PMF α) (f : α → β) (a : α) : p a ≤ (p.map f) (f a) := by
classical
rw [PMF.map_apply]
simpa using ENNReal.le_tsum (f := fun x => if f a = f x then p x else 0) a

/-- A randomized continuation preserves at least the joint mass of any one execution path. -/
theorem mul_le_bind_apply (p : PMF α) (f : α → PMF β) (a : α) (b : β) :
p a * f a b ≤ p.bind f b := by
rw [PMF.bind_apply]
exact ENNReal.le_tsum (f := fun x => p x * f x b) a

/-- The real-valued probabilities of a finite distribution sum to one. -/
theorem sum_toReal [Fintype α] (p : PMF α) :
∑ a, (p a).toReal = 1 := by simpa using tsum_toReal p
Expand Down Expand Up @@ -108,6 +137,47 @@ theorem toOuterMeasure_bind_toReal [Fintype α] (p : PMF α) (kernel : α → PM
(fun a _ => ENNReal.mul_ne_top (PMF.apply_ne_top _ _) (toOuterMeasure_ne_top _ _))]
simp

/-- A randomized continuation can fail because its precondition was already false or because
it fails from a valid input. No independence or finite-state assumption is needed. -/
theorem toOuterMeasure_bind_failure_le (p : PMF α) (kernel : α → PMF β)
(pre : Set α) (failure : Set β) (error : ENNReal)
(hstep : ∀ a ∈ p.support, a ∈ pre → (kernel a).toOuterMeasure failure ≤ error) :
(p.bind kernel).toOuterMeasure failure ≤ p.toOuterMeasure preᶜ + error := by
rw [PMF.toOuterMeasure_bind_apply, PMF.toOuterMeasure_apply, ← one_mul error, ← p.tsum_coe,
← ENNReal.tsum_mul_right, ← ENNReal.tsum_add]
refine ENNReal.tsum_le_tsum fun a => ?_
by_cases ha : a ∈ pre
· rw [Set.indicator_of_notMem (by simpa using ha), zero_add]
by_cases hs : a ∈ p.support
· exact mul_le_mul_of_nonneg_left (hstep a hs ha) zero_le
· simp [(PMF.apply_eq_zero_iff p a).2 hs]
· rw [Set.indicator_of_mem ha]
exact le_add_right (mul_le_of_le_one_right' (toOuterMeasure_le_one _ _))

/-- The real-valued error rule for a randomized continuation. -/
theorem toOuterMeasure_bind_failure_toReal_le (p : PMF α) (kernel : α → PMF β)
(pre : Set α) (failure : Set β) {error : ℝ} (herror : 0 ≤ error)
(hstep : ∀ a ∈ p.support, a ∈ pre →
((kernel a).toOuterMeasure failure).toReal ≤ error) :
((p.bind kernel).toOuterMeasure failure).toReal ≤
(p.toOuterMeasure preᶜ).toReal + error := by
rw [← ENNReal.toReal_ofReal herror, ← ENNReal.toReal_add (toOuterMeasure_ne_top _ _)
ofReal_ne_top]
exact ENNReal.toReal_mono (add_ne_top.2 ⟨toOuterMeasure_ne_top _ _, ofReal_ne_top⟩)
(toOuterMeasure_bind_failure_le p kernel pre failure _ fun a ha hpre =>
(le_ofReal_iff_toReal_le (toOuterMeasure_ne_top _ _) herror).2 (hstep a ha hpre))

/-- If a randomized test beats a threshold on average, at least its excess probability mass
consists of inputs whose conditional acceptance probability reaches that threshold. -/
theorem toOuterMeasure_rate_ge (p : PMF α) (kernel : α → PMF β) (event : Set β)
{threshold : ℝ} (hthreshold : 0 ≤ threshold) :
((p.bind kernel).toOuterMeasure event).toReal - threshold ≤
(p.toOuterMeasure {a | threshold ≤ ((kernel a).toOuterMeasure event).toReal}).toReal := by
have h := toOuterMeasure_bind_failure_toReal_le p kernel
{a | ((kernel a).toOuterMeasure event).toReal < threshold} event hthreshold
(fun a _ ha => le_of_lt ha)
simpa [Set.compl_ofPred] using h

/-- Randomized postprocessing averages the outcome probabilities over any discrete input. -/
theorem bind_apply_toReal_tsum (p : PMF α) (kernel : α → PMF β) (b : β) :
(p.bind kernel b).toReal = ∑' a, (p a).toReal * (kernel a b).toReal := by
Expand Down Expand Up @@ -162,6 +232,27 @@ theorem sum_mul_le [Fintype α] (p : PMF α) (score : α → ℝ) (bound : ℝ)
mul_le_mul_of_nonneg_left (hscore a) ENNReal.toReal_nonneg)
_ = bound := by rw [← Finset.sum_mul, sum_toReal, one_mul]

/-- The mass of an image of a uniform input is its fiber size divided by the input size. -/
theorem uniformOfFintype_map_apply [Fintype α] [Nonempty α] (f : α → β) (b : β) :
((PMF.uniformOfFintype α).map f) b =
(Nat.card {a // f a = b} : ℝ≥0∞) / Fintype.card α := by
classical
simp [← Finset.sum_filter, div_eq_mul_inv, Fintype.card_subtype, eq_comm]

/-- A continuation only needs to agree on outcomes which the preceding distribution can produce. -/
theorem bind_congr_on_support (p : PMF α) (f g : α → PMF β)
(h : ∀ a ∈ p.support, f a = g a) : p.bind f = p.bind g := by
ext b
refine tsum_congr fun a => ?_
by_cases ha : a ∈ p.support
· rw [h a ha]
· simp [(PMF.apply_eq_zero_iff p a).2 ha]

/-- Re-encoding only needs to agree on outcomes which the distribution can produce. -/
theorem map_congr_on_support (p : PMF α) (f g : α → β)
(h : ∀ a ∈ p.support, f a = g a) : p.map f = p.map g :=
bind_congr_on_support p _ _ (fun a ha => congrArg PMF.pure (h a ha))

/-- Evaluating the "pairing" bind `(do let a ← p; return (a, ← f a))` at `(a, b)`
gives the product `p a * f a b`. -/
theorem bind_pair_apply (p : PMF α) (f : α → PMF β) (a : α) (b : β) :
Expand All @@ -172,12 +263,39 @@ theorem bind_pair_apply (p : PMF α) (f : α → PMF β) (a : α) (b : β) :
· intro b' hb'; simp [PMF.pure_apply, hb'.symm]
· intro a' ha'; rw [PMF.bind_apply]; simp [PMF.pure_apply, ha'.symm]

/-- The pairing law also holds when the second sample's type depends on the first. -/
theorem bind_sigma_apply {β : α → Type*} (p : PMF α) (f : (a : α) → PMF (β a))
(a : α) (b : β a) :
(p.bind (fun a => (f a).map (Sigma.mk a))) ⟨a, b⟩ = p a * f a b := by
classical
rw [PMF.bind_apply, tsum_eq_single a]
· congr 1
simp
· intro i hi
have hne (c : β i) : (Sigma.mk a b : Sigma β) ≠ ⟨i, c⟩ :=
fun h => hi (congrArg Sigma.fst h).symm
simp [hne]

/-- Summing the pairing bind over the first component gives the marginal. -/
theorem bind_pair_tsum_fst (p : PMF α) (f : α → PMF β) (b : β) :
∑' a, (p.bind fun a' => (f a').bind fun b' => PMF.pure (a', b')) (a, b) =
(p.bind f) b := by
simp_rw [bind_pair_apply, PMF.bind_apply]

/-- Marginalizing a joint distribution over its freshly sampled second component. -/
@[simp] theorem map_fst_bind_pair (p : PMF α) (f : α → PMF β) :
(p.bind (fun a => (f a).map (a, ·))).map Prod.fst = p := by
simp only [PMF.map_bind, PMF.map_comp, Function.comp_def]
change p.bind (fun a => (f a).map (Function.const β a)) = p
simp

/-- Marginalizing a joint distribution over its first component gives ordinary sequencing. -/
@[simp] theorem map_snd_bind_pair (p : PMF α) (f : α → PMF β) :
(p.bind (fun a => (f a).map (a, ·))).map Prod.snd = p.bind f := by
simp only [PMF.map_bind, PMF.map_comp, Function.comp_def]
change p.bind (fun a => (f a).map id) = p.bind f
simp [PMF.map_id]

/-- A uniform distribution on a finite type is invariant under any equivalence. -/
theorem uniformOfFintype_map_equiv {γ : Type v} [Fintype α] [Fintype γ] [Nonempty α] [Nonempty γ]
(e : α ≃ γ) :
Expand All @@ -187,6 +305,15 @@ theorem uniformOfFintype_map_equiv {γ : Type v} [Fintype α] [Fintype γ] [None
· simp [Fintype.card_congr e]
· exact fun a ha => ite_eq_right fun h => ha (by simp [h])

/-- A uniform pair consists of two independent uniform samples. -/
theorem uniformOfFintype_prod [Fintype α] [Fintype β] [Nonempty α] [Nonempty β] :
PMF.uniformOfFintype (α × β) = (PMF.uniformOfFintype α).bind
(fun a => (PMF.uniformOfFintype β).map (fun b => (a, b))) := by
ext ⟨a, b⟩
simp only [PMF.map, Function.comp_def, bind_pair_apply, PMF.uniformOfFintype_apply,
Fintype.card_prod, Nat.cast_mul]
rw [ENNReal.mul_inv] <;> simp

/-- The posterior distribution `Pr[A = a | B = b]` as a `PMF`,
given `a ← p`, `b ← f a`, and that `b` has positive marginal probability:
the joint distribution's slice at `b`, normalized. -/
Expand Down
Loading