From 4fcddf5864976bd3404b514512ae85746bd5a32b Mon Sep 17 00:00:00 2001 From: Samuel Schlesinger Date: Sat, 3 Oct 2026 23:29:49 -0400 Subject: [PATCH 1/2] feat(Probability): bound statistical distance under scores, conditioning, and pairing Adds the standard ways statistical distance enters reductions, for finite types. Changing the distribution moves the expectation of a score bounded by `B` by at most `2 * B * dist p q` (`abs_sum_mul_sub_le`). Conditioning on an event changes a distribution by exactly the probability it discards (`dist_filter`). With a shared revealed input, statistical distance is the average distance of the conditional outputs (`dist_bind_pair`), and its square is at most the average squared distance (`dist_bind_pair_sq_le`). --- Cslib/Probability/StatisticalDistance.lean | 67 ++++++++++++++++++++++ 1 file changed, 67 insertions(+) diff --git a/Cslib/Probability/StatisticalDistance.lean b/Cslib/Probability/StatisticalDistance.lean index b27f61e15..3d8c4d98e 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,32 @@ 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, neg_sub] + · simp [ha] + 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 +217,25 @@ 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, + ENNReal.toReal_mul, ← mul_sub, abs_mul, abs_of_nonneg ENNReal.toReal_nonneg, + ← 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 := From 9a4128d2c81277ca31ee0dc3ffec24aea9fe1db8 Mon Sep 17 00:00:00 2001 From: Samuel Schlesinger Date: Sun, 4 Oct 2026 19:01:54 -0400 Subject: [PATCH 2/2] style(Probability): make a non-terminal simp rigid in statistical distance lemmas --- Cslib/Probability/StatisticalDistance.lean | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/Cslib/Probability/StatisticalDistance.lean b/Cslib/Probability/StatisticalDistance.lean index 3d8c4d98e..391983c32 100644 --- a/Cslib/Probability/StatisticalDistance.lean +++ b/Cslib/Probability/StatisticalDistance.lean @@ -148,8 +148,9 @@ theorem dist_filter [Finite α] (p : PMF α) (event : Set α) 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, neg_sub] - · simp [ha] + 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, @@ -223,9 +224,8 @@ theorem dist_bind_pair [Fintype α] [Finite β] (p : PMF α) (f g : α → PMF (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, - ENNReal.toReal_mul, ← mul_sub, abs_mul, abs_of_nonneg ENNReal.toReal_nonneg, - ← Finset.mul_sum, Finset.sum_div, mul_div_assoc] + 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 β) :