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
67 changes: 67 additions & 0 deletions Cslib/Probability/StatisticalDistance.lean
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,11 @@ deterministic case.
distance
- `dist_eq_one_of_disjoint_support`: PMFs with disjoint supports are at the
maximum statistical distance
- `abs_sum_mul_sub_le`: the expectation of a bounded score changes by at most twice the bound
times the statistical distance
- `dist_filter`: conditioning changes a finite distribution by exactly the discarded probability
- `dist_bind_pair`: with a shared revealed input, statistical distance is the average distance
of the conditional outputs
- `StatisticallyClose.trans`: closeness bounds chain through an intermediate
distribution, adding the errors
- `statisticallyClose_zero_iff`: zero error is equality
Expand Down Expand Up @@ -99,6 +104,23 @@ theorem dist_eq [Fintype α] (p q : PMF α) :
dist p q = (∑ a, |(p a).toReal - (q a).toReal|) / 2 := by
simp [dist_eq_tsum]

/-- Changing the distribution changes the expectation of a bounded real-valued score by at most
twice its absolute bound times the statistical distance. Scores with values in `[-1, 1]` use
`bound = 1`. -/
theorem abs_sum_mul_sub_le [Fintype α] (p q : PMF α) (score : α → ℝ)
{bound : ℝ} (hscore : ∀ a, |score a| ≤ bound) :
|(∑ a, (p a).toReal * score a) - ∑ a, (q a).toReal * score a| ≤
2 * bound * dist p q := by
rw [← Finset.sum_sub_distrib]
simp only [← sub_mul]
calc
_ ≤ ∑ a, |((p a).toReal - (q a).toReal) * score a| := Finset.abs_sum_le_sum_abs _ _
_ ≤ ∑ a, |(p a).toReal - (q a).toReal| * bound := by
simp only [abs_mul]
gcongr with a
exact hscore a
_ = _ := by rw [← Finset.sum_mul, dist_eq]; ring

/-- Statistical distance is at most one. -/
theorem dist_le_one (p q : PMF α) : dist p q ≤ 1 := by
rw [dist_eq_tsum]
Expand All @@ -108,6 +130,33 @@ theorem dist_le_one (p q : PMF α) : dist p q ≤ 1 := by
rw [(summable_toReal p).tsum_add (summable_toReal q), tsum_toReal, tsum_toReal] at h
linarith

/-- Conditioning on a possible event changes the distribution by exactly the probability
discarded outside that event. -/
theorem dist_filter [Finite α] (p : PMF α) (event : Set α)
(hevent : ∃ a ∈ event, a ∈ p.support) :
dist p (p.filter event hevent) = (p.toOuterMeasure eventᶜ).toReal := by
classical
let := Fintype.ofFinite α
have hpositive := (toOuterMeasure_toReal_pos_iff p event).mpr hevent
have hprobability : (p.toOuterMeasure event).toReal ≤ 1 := by
simpa using ENNReal.toReal_mono ENNReal.one_ne_top (toOuterMeasure_le_one p event)
have habsolute (a : α) : |(p a).toReal - ((p.filter event hevent) a).toReal| =
((p.filter event hevent) a).toReal - (p a).toReal +
2 * (if a ∈ eventᶜ then (p a).toReal else 0) := by
by_cases ha : a ∈ event
· have hle : (p a).toReal ≤ ((p.filter event hevent) a).toReal := by
rw [filter_apply_toReal, ite_eq_left ha]
exact le_div_self ENNReal.toReal_nonneg hpositive hprobability
rw [abs_of_nonpos (sub_nonpos.mpr hle)]
simp [ha]
· simp only [filter_apply_toReal, ha, Set.mem_compl_iff, not_false_eq_true, ↓reduceIte,
sub_zero, ENNReal.abs_toReal]
ring
rw [dist_eq, toOuterMeasure_apply_toReal]
simp only [habsolute, Finset.sum_add_distrib, Finset.sum_sub_distrib, sum_toReal, sub_self,
zero_add, ← Finset.mul_sum]
ring_nf

/-- PMFs with disjoint supports are at the maximum statistical distance. -/
theorem dist_eq_one_of_disjoint_support {p q : PMF α}
(h : Disjoint p.support q.support) : dist p q = 1 := by
Expand Down Expand Up @@ -169,6 +218,24 @@ theorem dist_map_le (p q : PMF α) (f : α → β) :
dist (p.map f) (q.map f) ≤ dist p q := by
simpa [PMF.bind_pure_comp] using dist_bind_le p q (PMF.pure ∘ f)

/-- With a shared revealed input, distance is the average distance of the conditional outputs. -/
theorem dist_bind_pair [Fintype α] [Finite β] (p : PMF α) (f g : α → PMF β) :
dist (p.bind (fun a => (f a).map (a, ·)))
(p.bind (fun a => (g a).map (a, ·))) =
∑ a, (p a).toReal * dist (f a) (g a) := by
let := Fintype.ofFinite β
simp only [dist_eq, Fintype.sum_prod_type, PMF.map, Function.comp_def, bind_pair_apply]
simp [← mul_sub, ← Finset.mul_sum, Finset.sum_div, mul_div_assoc]

/-- Squared distance with a shared revealed input is at most the average squared distance. -/
theorem dist_bind_pair_sq_le [Fintype α] [Finite β] (p : PMF α) (f g : α → PMF β) :
dist (p.bind (fun a => (f a).map (a, ·)))
(p.bind (fun a => (g a).map (a, ·))) ^ 2 ≤
∑ a, (p a).toReal * dist (f a) (g a) ^ 2 := by
rw [dist_bind_pair]
exact Real.pow_arith_mean_le_arith_mean_pow _ _ _ (fun _ _ => ENNReal.toReal_nonneg)
(sum_toReal p) (fun _ _ => dist_nonneg) 2

/-- Two PMFs are `ε`-statistically close when their statistical distance is at
most `ε`. The `ℝ≥0` parameter rules out meaningless negative bounds. -/
def StatisticallyClose (p q : PMF α) (ε : ℝ≥0) : Prop :=
Expand Down
Loading