diff --git a/Cslib/Probability/StatisticalDistance.lean b/Cslib/Probability/StatisticalDistance.lean index b27f61e15..391983c32 100644 --- a/Cslib/Probability/StatisticalDistance.lean +++ b/Cslib/Probability/StatisticalDistance.lean @@ -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 @@ -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] @@ -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 @@ -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 :=