diff --git a/SampCert/DifferentialPrivacy/Pure/Mechanism/Code.lean b/SampCert/DifferentialPrivacy/Pure/Mechanism/Code.lean index 6d82adf7..d9ada6e8 100644 --- a/SampCert/DifferentialPrivacy/Pure/Mechanism/Code.lean +++ b/SampCert/DifferentialPrivacy/Pure/Mechanism/Code.lean @@ -20,3 +20,5 @@ def privNoisedQueryPure (query : List T → ℤ) (Δ : ℕ+) (ε₁ ε₂ : ℕ+ DiscreteLaplaceGenSamplePMF (Δ * ε₂) ε₁ (query l) end SLang + +end diff --git a/SampCert/DifferentialPrivacy/Queries/MWEM/Basic.lean b/SampCert/DifferentialPrivacy/Queries/MWEM/Basic.lean new file mode 100644 index 00000000..d83ef990 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/MWEM/Basic.lean @@ -0,0 +1,8 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.MWEM.Code +import SampCert.DifferentialPrivacy.Queries.MWEM.Properties diff --git a/SampCert/DifferentialPrivacy/Queries/MWEM/Code.lean b/SampCert/DifferentialPrivacy/Queries/MWEM/Code.lean new file mode 100644 index 00000000..a4e7b35c --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/MWEM/Code.lean @@ -0,0 +1,102 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Abstract +import SampCert.DifferentialPrivacy.Pure.System +import SampCert.DifferentialPrivacy.Queries.ReportNoisyMax.Code + +/-! +# MWEM mechanism (Hardt 2012) + +Each round selects a query via Report-Noisy-Max, releases a noisy answer via the +Laplace mechanism, and post-processes the synthetic state. The synthetic state +type and update rule are abstracted as a `SyntheticUpdater`, so the privacy +proof works uniformly over any choice of updater (real-valued multiplicative +weights, rational, integer, identity, etc.). +-/ + +noncomputable section + +namespace SLang + +variable {X : Type} + +structure SyntheticUpdater (X : Type) (n : ℕ) (Δ : ℕ+) (State : Type) (init : State) where + queries : Fin (n+1) → List X → ℤ + scoreFn : State → Fin (n+1) → List X → ℤ + update : State → Fin (n+1) → ℤ → State + queries_sens : ∀ i, sensitivity (queries i) Δ + scoreFn_sens : ∀ A i, sensitivity (scoreFn A i) Δ + +variable {n : ℕ} {Δ : ℕ+} {State : Type} {init : State} + +def mwemRound (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (A : State) : + Mechanism X (Fin (n+1) × ℤ) := + privComposeAdaptive + (privReportNoisyMax n (U.scoreFn A) Δ ε₁ ε₂) + (fun i => privNoisedQueryPure (U.queries i) Δ ε₁ ε₂) + +def mwem (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) : + ℕ → State → Mechanism X (List (Fin (n+1) × ℤ)) + | 0, _ => privConst [] + | T+1, A => + privPostProcess + (privComposeAdaptive (mwemRound U ε₁ ε₂ A) + (fun im => mwem U ε₁ ε₂ T (U.update A im.1 im.2))) + (fun (im, hist) => im :: hist) + +end SLang + +end + +namespace SLang + +variable {X : Type} {n : ℕ} {Δ : ℕ+} {State : Type} {init : State} + +def mwemRoundSLang (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (A : State) + (l : List X) : SLang (Fin (n+1) × ℤ) := do + let i ← privReportNoisyMaxSLang n (U.scoreFn A) Δ ε₁ ε₂ l + let m ← DiscreteLaplaceGenSample (Δ * ε₂) ε₁ (U.queries i l) + return (i, m) + +def mwemSLang (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) : + ℕ → State → List X → SLang (List (Fin (n+1) × ℤ)) + | 0, _, _ => probPure [] + | T+1, A, l => + probBind + (probBind (mwemRoundSLang U ε₁ ε₂ A l) (fun im => + probBind (mwemSLang U ε₁ ε₂ T (U.update A im.1 im.2) l) (fun hist => + probPure (im, hist)))) + (fun p => probPure (p.1 :: p.2)) + +theorem mwemRoundSLang_eq (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (A : State) + (l : List X) : + mwemRoundSLang U ε₁ ε₂ A l = ((mwemRound U ε₁ ε₂ A l : SPMF _) : SLang _) := by + show probBind (privReportNoisyMaxSLang n (U.scoreFn A) Δ ε₁ ε₂ l) _ = _ + rw [privReportNoisyMaxSLang_eq] + rfl + +theorem mwemSLang_eq (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) : + ∀ (T : ℕ) (A : State) (l : List X), + mwemSLang U ε₁ ε₂ T A l = ((mwem U ε₁ ε₂ T A l : SPMF _) : SLang _) := by + intro T + induction T with + | zero => intro A l; rfl + | succ T IH => + intro A l + show probBind (probBind (mwemRoundSLang U ε₁ ε₂ A l) _) _ = _ + rw [mwemRoundSLang_eq] + show probBind (probBind _ + (fun im : Fin (n+1) × ℤ => + probBind (mwemSLang U ε₁ ε₂ T (U.update A im.1 im.2) l) _)) _ = _ + conv_lhs => enter [1, 2, im]; rw [IH] + rfl + +def mwemSPMF (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (T : ℕ) (A : State) + (l : List X) : SPMF (List (Fin (n+1) × ℤ)) := + ⟨ mwemSLang U ε₁ ε₂ T A l, mwemSLang_eq U ε₁ ε₂ T A l ▸ (mwem U ε₁ ε₂ T A l).2 ⟩ + +end SLang diff --git a/SampCert/DifferentialPrivacy/Queries/MWEM/Properties.lean b/SampCert/DifferentialPrivacy/Queries/MWEM/Properties.lean new file mode 100644 index 00000000..f5833296 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/MWEM/Properties.lean @@ -0,0 +1,68 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.MWEM.Code +import SampCert.DifferentialPrivacy.Queries.ReportNoisyMax.Basic +import SampCert.DifferentialPrivacy.Pure.System + +/-! +# Privacy of MWEM + +Each round of MWEM is `2(ε₁/ε₂)`-DP by adaptive composition of Report-Noisy-Max +and the Laplace mechanism, both `(ε₁/ε₂)`-DP under the sensitivity hypotheses +of `SyntheticUpdater`. The synthetic state update is post-processing of the +selection-and-measurement transcript, so it adds no privacy cost. By induction +over `T` rounds, MWEM is `2T(ε₁/ε₂)`-DP. +-/ + +noncomputable section + +open Classical Nat Int Real ENNReal + +namespace SLang + +variable {X : Type} {n : ℕ} {Δ : ℕ+} {State : Type} {init : State} + +instance : MeasurableSpace (Fin (n+1) × ℤ) := ⊤ +instance : DiscreteMeasurableSpace (Fin (n+1) × ℤ) where + forall_measurableSet _ := .congr trivial rfl + +instance : MeasurableSpace (List (Fin (n+1) × ℤ)) := ⊤ +instance : DiscreteMeasurableSpace (List (Fin (n+1) × ℤ)) where + forall_measurableSet _ := .congr trivial rfl + +theorem mwemRound_DP (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (ε : NNReal) + (HN : laplace_pureDP_noise_priv ε₁ ε₂ ε) (A : State) : + PureDPSystem.prop (mwemRound U ε₁ ε₂ A) (ε + ε) := + PureDPSystem.adaptive_compose_prop + (privReportNoisyMax_DP n (U.scoreFn A) Δ ε₁ ε₂ ε HN (U.scoreFn_sens A)) + (fun i => privNoisedQueryPure_DP (U.queries i) Δ ε₁ ε₂ ε HN (U.queries_sens i)) + rfl + +theorem mwem_DP (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (ε : NNReal) + (HN : laplace_pureDP_noise_priv ε₁ ε₂ ε) (T : ℕ) (A : State) : + PureDPSystem.prop (mwem U ε₁ ε₂ T A) (T * (ε + ε)) := by + induction T generalizing A with + | zero => exact PureDPSystem.const_prop (by push_cast; ring) + | succ T IH => + refine PureDPSystem.postprocess_prop <| + PureDPSystem.adaptive_compose_prop (mwemRound_DP U ε₁ ε₂ ε HN A) + (fun im => IH (U.update A im.1 im.2)) ?_ + push_cast; ring + +theorem mwemSPMF_DP (U : SyntheticUpdater X n Δ State init) (ε₁ ε₂ : ℕ+) (ε : NNReal) + (HN : laplace_pureDP_noise_priv ε₁ ε₂ ε) (T : ℕ) (A : State) : + PureDPSystem.prop (mwemSPMF U ε₁ ε₂ T A) (T * (ε + ε)) := by + have heq : (fun l => mwemSPMF U ε₁ ε₂ T A l) = (fun l => (mwem U ε₁ ε₂ T A l : SPMF _)) := by + funext l + apply Subtype.ext + exact mwemSLang_eq U ε₁ ε₂ T A l + rw [show mwemSPMF U ε₁ ε₂ T A = _ from heq] + exact mwem_DP U ε₁ ε₂ ε HN T A + +end SLang + +end diff --git a/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Float.lean b/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Float.lean new file mode 100644 index 00000000..b592e011 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Float.lean @@ -0,0 +1,103 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.MWEM.Code + +/-! +# Float-valued multiplicative-weights updater + +Hardt's MWEM (Hardt, Ligett, McSherry 2012, Algorithm 1) with the synthetic +state implemented in IEEE-754 doubles. This is the representation used by +deployed systems (SmartNoise, etc.). + +The synthetic state lives in `Fin numBins → Float`. Following Figure 1 of the +paper: + +* `A_0` is the uniform distribution scaled by `n` (dataset size), so each bin + starts with mass `n / numBins` and the total is `n`. +* The update rule for each round `i` is + `A_i(x) ∝ A_{i-1}(x) · exp( q_i(x) · (m_i - q_i(A_{i-1})) / (2n) )` + with renormalization to fixed total `n`. + +Privacy comes from the Laplace mechanism in the measurement step, which uses +exact integer arithmetic. The Float update is post-processing of the integer +transcript and contributes no privacy cost. Lean's Float type has essentially +no axiomatic theory, but the privacy proof requires only that the update be a +total function of its arguments. + +The score function uses `Float.toInt32` to produce an integer constant (in `D`) +for sensitivity reasons; its specific value depends on IEEE-754 semantics but +is irrelevant to the sensitivity bound. +-/ + +open Classical Nat Int Real ENNReal + +namespace SLang + +variable {numBins : ℕ+} {n : ℕ} {Δ : ℕ+} + +def intToFloat (n : ℤ) : Float := + match n with + | .ofNat k => k.toUInt64.toFloat + | .negSucc k => -((k + 1).toUInt64.toFloat) + +def floatToInt (x : Float) : ℤ := x.toInt32.toInt + +def floatLinearQuery (q : Fin numBins → ℤ) (D : List (Fin numBins)) : ℤ := + (D.map q).sum + +def floatLinearQueryFloat (q : Fin numBins → ℤ) (A : Fin numBins → Float) : Float := + ((List.finRange numBins).map (fun b => intToFloat (q b) * A b)).foldr (· + ·) 0.0 + +def floatScore (q : Fin numBins → ℤ) (A : Fin numBins → Float) (D : List (Fin numBins)) : ℤ := + |floatLinearQuery q D - floatToInt (floatLinearQueryFloat q A)| + +def floatMWUpdate (n : ℕ) (q : Fin numBins → ℤ) (m : ℤ) (A : Fin numBins → Float) : + Fin numBins → Float := + let nF : Float := intToFloat (n : ℤ) + let err : Float := (intToFloat m - floatLinearQueryFloat q A) / (2.0 * nF) + let unnorm : Fin numBins → Float := fun b => A b * Float.exp (intToFloat (q b) * err) + let total : Float := + ((List.finRange numBins).map unnorm).foldr (· + ·) 0.0 + fun b => unnorm b * nF / total + +def floatInit (n : ℕ) (numBins : ℕ+) : Fin numBins → Float := + fun _ => intToFloat (n : ℤ) / intToFloat (numBins.val : ℤ) + +theorem floatLinearQuery_sens (q : Fin numBins → ℤ) + (Hbound : ∀ b, |q b| ≤ (Δ : ℤ)) : sensitivity (floatLinearQuery q) Δ := by + intro l₁ l₂ Hn + have qnat : ∀ k : Fin numBins, (q k).natAbs ≤ (Δ : ℕ) := fun k => by + have := Hbound k; rw [Int.abs_eq_natAbs] at this; exact_mod_cast this + cases Hn with + | @Addition a b k h1 h2 => subst h1 h2; simp [floatLinearQuery]; exact qnat _ + | @Deletion a b k h1 h2 => subst h1 h2; simp [floatLinearQuery]; exact qnat _ + +lemma natAbs_abs_sub_abs_le' {a b : ℤ} : (|a| - |b|).natAbs ≤ (a - b).natAbs := by + zify; exact abs_abs_sub_abs_le_abs_sub a b + +theorem floatScore_sens (q : Fin numBins → ℤ) (A : Fin numBins → Float) + (Hbound : ∀ b, |q b| ≤ (Δ : ℤ)) : sensitivity (floatScore q A) Δ := by + intro l₁ l₂ Hn + have hQ : (floatLinearQuery q l₁ - floatLinearQuery q l₂).natAbs ≤ (Δ : ℕ) := + floatLinearQuery_sens (Δ := Δ) q Hbound l₁ l₂ Hn + set c : ℤ := floatToInt (floatLinearQueryFloat q A) + show (|floatLinearQuery q l₁ - c| - |floatLinearQuery q l₂ - c|).natAbs ≤ (Δ : ℕ) + refine le_trans natAbs_abs_sub_abs_le' ?_ + rw [show floatLinearQuery q l₁ - c - (floatLinearQuery q l₂ - c) = + floatLinearQuery q l₁ - floatLinearQuery q l₂ from by ring] + exact hQ + +def floatMWUpdater (nData : ℕ) (q : Fin (n+1) → Fin numBins → ℤ) + (Hbound : ∀ i b, |q i b| ≤ (Δ : ℤ)) : + SyntheticUpdater (Fin numBins) n Δ (Fin numBins → Float) (floatInit nData numBins) where + queries i := floatLinearQuery (q i) + scoreFn A i := floatScore (q i) A + update A i m := floatMWUpdate nData (q i) m A + queries_sens i := floatLinearQuery_sens (q i) (Hbound i) + scoreFn_sens A i := floatScore_sens (q i) A (Hbound i) + +end SLang diff --git a/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Real.lean b/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Real.lean new file mode 100644 index 00000000..7f31f446 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/MWEM/Updaters/Real.lean @@ -0,0 +1,90 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.MWEM.Code + +/-! +# Real-valued multiplicative-weights updater + +Hardt's MWEM (Hardt, Ligett, McSherry 2012, Algorithm 1) with the synthetic +state in `Fin numBins → NNReal`. This is the textbook reference instantiation; +not extractable to executable code (uses `Real.exp` and real arithmetic). + +Following Figure 1 of the paper: + +* `A_0(x) = n / numBins` for all `x` (uniform with total mass `n`). +* The update for round `i` is + `A_i(x) ∝ A_{i-1}(x) · exp( q_i(x) · (m_i - q_i(A_{i-1})) / (2n) )` + with renormalization to fixed total `n`. + +If the renormalization total degenerates to zero, the updater returns the zero +histogram. The privacy proof does not depend on non-degeneracy. +-/ + +noncomputable section + +open Classical Nat Int Real ENNReal NNReal + +namespace SLang + +variable {numBins : ℕ+} {n : ℕ} {Δ : ℕ+} + +def realLinearQuery (q : Fin numBins → ℤ) (D : List (Fin numBins)) : ℤ := + (D.map q).sum + +def realLinearQueryReal (q : Fin numBins → ℤ) (A : Fin numBins → NNReal) : ℝ := + ∑ b, (q b : ℝ) * (A b : ℝ) + +def realScore (q : Fin numBins → ℤ) (A : Fin numBins → NNReal) (D : List (Fin numBins)) : ℤ := + |realLinearQuery q D - ⌊realLinearQueryReal q A⌋| + +def realMWUpdate (n : ℕ) (q : Fin numBins → ℤ) (m : ℤ) (A : Fin numBins → NNReal) : + Fin numBins → NNReal := + let nR : ℝ := (n : ℝ) + let err : ℝ := ((m : ℝ) - realLinearQueryReal q A) / (2 * nR) + let unnorm : Fin numBins → ℝ := fun b => (A b : ℝ) * Real.exp ((q b : ℝ) * err) + let total : ℝ := ∑ b, unnorm b + fun b => Real.toNNReal (unnorm b * nR / total) + +def realInit (n : ℕ) (numBins : ℕ+) : Fin numBins → NNReal := + fun _ => Real.toNNReal ((n : ℝ) / (numBins.val : ℝ)) + +theorem realLinearQuery_sens (q : Fin numBins → ℤ) + (Hbound : ∀ b, |q b| ≤ (Δ : ℤ)) : sensitivity (realLinearQuery q) Δ := by + intro l₁ l₂ Hn + have qnat : ∀ k : Fin numBins, (q k).natAbs ≤ (Δ : ℕ) := fun k => by + have := Hbound k; rw [Int.abs_eq_natAbs] at this; exact_mod_cast this + cases Hn with + | @Addition a b k h1 h2 => subst h1 h2; simp [realLinearQuery]; exact qnat _ + | @Deletion a b k h1 h2 => subst h1 h2; simp [realLinearQuery]; exact qnat _ + +lemma natAbs_abs_sub_abs_le {a b : ℤ} : (|a| - |b|).natAbs ≤ (a - b).natAbs := by + zify; exact abs_abs_sub_abs_le_abs_sub a b + +theorem realScore_sens (q : Fin numBins → ℤ) (A : Fin numBins → NNReal) + (Hbound : ∀ b, |q b| ≤ (Δ : ℤ)) : sensitivity (realScore q A) Δ := by + intro l₁ l₂ Hn + have hQ : (realLinearQuery q l₁ - realLinearQuery q l₂).natAbs ≤ (Δ : ℕ) := + realLinearQuery_sens (Δ := Δ) q Hbound l₁ l₂ Hn + set c : ℤ := ⌊realLinearQueryReal q A⌋ + show (|realLinearQuery q l₁ - c| - |realLinearQuery q l₂ - c|).natAbs ≤ (Δ : ℕ) + refine le_trans natAbs_abs_sub_abs_le ?_ + rw [show realLinearQuery q l₁ - c - (realLinearQuery q l₂ - c) = + realLinearQuery q l₁ - realLinearQuery q l₂ from by ring] + exact hQ + +def realMWUpdater (nData : ℕ) (q : Fin (n+1) → Fin numBins → ℤ) + (Hbound : ∀ i b, |q i b| ≤ (Δ : ℤ)) : + SyntheticUpdater (Fin numBins) n Δ (Fin numBins → NNReal) (realInit nData numBins) where + queries i := realLinearQuery (q i) + scoreFn A i := realScore (q i) A + update A i m := realMWUpdate nData (q i) m A + queries_sens i := realLinearQuery_sens (q i) (Hbound i) + scoreFn_sens A i := realScore_sens (q i) A (Hbound i) + +end SLang + +end diff --git a/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Basic.lean b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Basic.lean new file mode 100644 index 00000000..d786cf62 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Basic.lean @@ -0,0 +1,8 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.ReportNoisyMax.Code +import SampCert.DifferentialPrivacy.Queries.ReportNoisyMax.Properties diff --git a/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Code.lean b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Code.lean new file mode 100644 index 00000000..6eeb5643 --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Code.lean @@ -0,0 +1,114 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Abstract +import SampCert.DifferentialPrivacy.Pure.System + +/-! +# Report Noisy Max (Laplace variant) + +Selection mechanism: given a finite, nonempty family of integer-valued queries +each of sensitivity `Δ`, draw independent discrete Laplace noise per query and +return the index of the largest noised score. This is `(ε₁/ε₂)`-DP independent +of the number of candidates. +-/ + +noncomputable section + +namespace SLang + +variable {T : Type} + +/-- +Sample independent Laplace noise for each of `n+1` queries, returning the joint +vector of noised values. Implemented as a fold over `Fin (n+1)`. +-/ +def privNoisedFamily (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) (l : List T) : + PMF (Fin (n+1) → ℤ) := + match n with + | 0 => + (privNoisedQueryPure (s 0) (2 * Δ) ε₁ ε₂ l).bind (fun z => PMF.pure (fun _ => z)) + | Nat.succ n' => + (privNoisedFamily n' (fun i => s i.castSucc) Δ ε₁ ε₂ l).bind (fun rest => + (privNoisedQueryPure (s (Fin.last (n'+1))) (2 * Δ) ε₁ ε₂ l).bind (fun last => + PMF.pure (fun i => if h : i.val < n'+1 then rest ⟨i.val, h⟩ else last))) + +/-- +Argmax of a function `Fin (n+1) → ℤ`, returning the largest `i` achieving the +maximum (ties broken by taking the larger index). +-/ +def Fin.argmax (n : ℕ) (f : Fin (n+1) → ℤ) : Fin (n+1) := + match n with + | 0 => 0 + | Nat.succ n' => + let prev := Fin.argmax n' (fun i => f i.castSucc) + let last : Fin (n'+2) := Fin.last (n'+1) + if f last < f prev.castSucc + then prev.castSucc + else last + +/-- +Report-Noisy-Max with discrete Laplace noise. + +Given `n+1` candidate queries `s : Fin (n+1) → List T → ℤ` each of sensitivity +`Δ`, sample independent Laplace noise (scaled to `(2Δ * ε₂) / ε₁`) for each, and +return the index of the largest noised value. + +The mechanism is `(ε₁/ε₂)`-DP regardless of `n`. +-/ +def privReportNoisyMax (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) (l : List T) : + PMF (Fin (n+1)) := + (privNoisedFamily n s Δ ε₁ ε₂ l).bind (fun noised => PMF.pure (Fin.argmax n noised)) + +end SLang + +end + +namespace SLang + +def privNoisedFamilySLang {T : Type} (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (l : List T) : SLang (Fin (n+1) → ℤ) := + match n with + | 0 => do + let z ← DiscreteLaplaceGenSample (2 * Δ * ε₂) ε₁ (s 0 l) + return (fun _ => z) + | Nat.succ n' => do + let rest ← privNoisedFamilySLang n' (fun i => s i.castSucc) Δ ε₁ ε₂ l + let last ← DiscreteLaplaceGenSample (2 * Δ * ε₂) ε₁ (s (Fin.last (n'+1)) l) + return (fun i => if h : i.val < n'+1 then rest ⟨i.val, h⟩ else last) + +def privReportNoisyMaxSLang {T : Type} (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (l : List T) : SLang (Fin (n+1)) := do + let noised ← privNoisedFamilySLang n s Δ ε₁ ε₂ l + return Fin.argmax n noised + +theorem privNoisedFamilySLang_eq {T : Type} (n : ℕ) (s : Fin (n+1) → List T → ℤ) + (Δ ε₁ ε₂ : ℕ+) (l : List T) : + privNoisedFamilySLang n s Δ ε₁ ε₂ l = ((privNoisedFamily n s Δ ε₁ ε₂ l : PMF _) : SLang _) := by + induction n with + | zero => + show probBind (DiscreteLaplaceGenSample (2 * Δ * ε₂) ε₁ (s 0 l)) + (fun z => probPure (fun _ => z)) = _ + rfl + | succ n IH => + show probBind (privNoisedFamilySLang n _ Δ ε₁ ε₂ l) _ = _ + rw [IH] + rfl + +theorem privReportNoisyMaxSLang_eq {T : Type} (n : ℕ) (s : Fin (n+1) → List T → ℤ) + (Δ ε₁ ε₂ : ℕ+) (l : List T) : + privReportNoisyMaxSLang n s Δ ε₁ ε₂ l = + ((privReportNoisyMax n s Δ ε₁ ε₂ l : PMF _) : SLang _) := by + show probBind (privNoisedFamilySLang n s Δ ε₁ ε₂ l) _ = _ + rw [privNoisedFamilySLang_eq] + rfl + +def privReportNoisyMaxSPMF {T : Type} (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (l : List T) : SPMF (Fin (n+1)) := + ⟨ privReportNoisyMaxSLang n s Δ ε₁ ε₂ l, + privReportNoisyMaxSLang_eq n s Δ ε₁ ε₂ l ▸ (privReportNoisyMax n s Δ ε₁ ε₂ l).2 ⟩ + +end SLang diff --git a/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Properties.lean b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Properties.lean new file mode 100644 index 00000000..39271a0b --- /dev/null +++ b/SampCert/DifferentialPrivacy/Queries/ReportNoisyMax/Properties.lean @@ -0,0 +1,276 @@ +/- +Copyright (c) 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Markus de Medeiros +-/ + +import SampCert.DifferentialPrivacy.Queries.ReportNoisyMax.Code +import SampCert.DifferentialPrivacy.Sensitivity +import SampCert.DifferentialPrivacy.Abstract +import SampCert.DifferentialPrivacy.Pure.System +import SampCert.DifferentialPrivacy.Queries.AboveThresh.Privacy + +/-! +# Privacy of `privReportNoisyMax` + +This file proves that the Report-Noisy-Max mechanism with discrete Laplace noise +is `(ε₁/ε₂)`-pure-DP, regardless of the number of candidates. + +The proof follows Dwork–Roth Claim 3.9 (adapted from Exponential to Laplace +noise): condition on the noises of all losers, reduce to a one-dimensional +Laplace tail-shift on the winner's noise, and observe that swapping `D` for a +neighbour `D'` shifts the winner's threshold by at most `2Δ`. +-/ + +noncomputable section + +open Classical Nat Int Real ENNReal + +namespace SLang + +variable {T : Type} + +lemma Fin.argmax_succ (n' : ℕ) (f : Fin (n'+2) → ℤ) : + Fin.argmax (n'+1) f = + (let prev := Fin.argmax n' (fun i => f i.castSucc) + if f (Fin.last (n'+1)) < f prev.castSucc then prev.castSucc else Fin.last (n'+1)) := + rfl + +lemma Fin.lt_succ_or_last {n : ℕ} (j : Fin (n+2)) : + (∃ k : Fin (n+1), j = k.castSucc) ∨ j = Fin.last (n+1) := by + rcases Nat.lt_succ_iff_lt_or_eq.mp j.isLt with h | h + · exact Or.inl ⟨⟨j.val, h⟩, Fin.ext rfl⟩ + · exact Or.inr (Fin.ext (by simp [h])) + +lemma Fin.argmax_eq_iff (n : ℕ) (f : Fin (n+1) → ℤ) (i : Fin (n+1)) : + Fin.argmax n f = i ↔ + (∀ j : Fin (n+1), j.val > i.val → f j < f i) ∧ + (∀ j : Fin (n+1), j.val < i.val → f j ≤ f i) := by + induction n with + | zero => + have hi : i = 0 := Fin.ext (Nat.lt_one_iff.mp i.isLt) + refine ⟨fun _ => ⟨fun j hj => ?_, fun j hj => ?_⟩, fun _ => ?_⟩ + · have := j.isLt; rw [hi] at hj; omega + · have := j.isLt; rw [hi] at hj; omega + · simp [Fin.argmax, hi] + | succ n IH => + rw [Fin.argmax_succ] + dsimp only + set prev := Fin.argmax n (fun j => f j.castSucc) + have IHprev := (IH (fun j => f j.castSucc) prev).mp rfl + refine ⟨fun Heq => ?_, fun ⟨H1, H2⟩ => ?_⟩ + · split_ifs at Heq with hlt + all_goals subst Heq + · refine ⟨fun j Hj => ?_, fun j Hj => ?_⟩ + · rcases Fin.lt_succ_or_last j with ⟨k, rfl⟩ | rfl + · exact IHprev.1 k Hj + · simpa using hlt + · rcases Fin.lt_succ_or_last j with ⟨k, rfl⟩ | rfl + · exact IHprev.2 k Hj + · simp at Hj; exact absurd prev.isLt (by omega) + · push Not at hlt + refine ⟨fun j Hj => ?_, fun j Hj => ?_⟩ + · have := j.isLt; simp at Hj; omega + · rcases Fin.lt_succ_or_last j with ⟨k, rfl⟩ | rfl + · refine le_trans ?_ hlt + rcases lt_trichotomy k.val prev.val with h | h | h + · exact IHprev.2 _ h + · rw [show k = prev from Fin.ext h] + · exact (IHprev.1 _ h).le + · simp at Hj + · split_ifs with hlt + · refine Fin.ext ?_ + by_contra hne + rcases lt_or_gt_of_ne hne with h | h + · rcases Fin.lt_succ_or_last i with ⟨k, rfl⟩ | rfl + · linarith [IHprev.1 k h, H2 prev.castSucc h] + · linarith [H2 prev.castSucc h] + · rcases Fin.lt_succ_or_last i with ⟨k, rfl⟩ | rfl + · linarith [IHprev.2 k h, H1 prev.castSucc h] + · exact absurd h (by have := prev.isLt; simp) + · push Not at hlt + refine Fin.ext ?_ + by_contra hne + rcases Fin.lt_succ_or_last i with ⟨k, rfl⟩ | rfl + · have hilast : (Fin.last (n+1)).val > k.castSucc.val := by simp; omega + have hH1 := H1 (Fin.last (n+1)) hilast + by_cases heq : k = prev + · rw [heq] at hH1; linarith + · rcases lt_or_gt_of_ne (fun h => heq (Fin.ext h)) with h | h + · linarith [IHprev.2 _ h] + · linarith [IHprev.1 _ h] + · exact hne rfl + +lemma fin_eq_iff_castSucc_last {n : ℕ} (f : Fin (n+2) → ℤ) (rest : Fin (n+1) → ℤ) (last : ℤ) : + f = (fun (i : Fin (n+2)) => if h : i.val < n+1 then rest ⟨i.val, h⟩ else last) ↔ + (∀ j : Fin (n+1), f j.castSucc = rest j) ∧ f (Fin.last (n+1)) = last := by + refine ⟨fun h => ?_, fun ⟨h_rest, h_last⟩ => ?_⟩ + · refine ⟨fun j => ?_, ?_⟩ + · rw [h]; simp [j.isLt] + · rw [h]; simp + · funext i + by_cases hi : i.val < n+1 + · rw [dif_pos hi, ← h_rest ⟨i.val, hi⟩]; rfl + · rw [dif_neg hi, ← h_last] + congr 1 + exact Fin.ext (by have := i.isLt; simp; omega) + +lemma privNoisedFamily_apply : ∀ (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (l : List T) (f : Fin (n+1) → ℤ), + (privNoisedFamily n s Δ ε₁ ε₂ l) f = + ∏ i : Fin (n+1), (privNoisedQueryPure (s i) (2 * Δ) ε₁ ε₂ l) (f i) := by + intro n + induction n with + | zero => + intro s Δ ε₁ ε₂ l f + rw [privNoisedFamily, PMF.bind_apply, tsum_eq_single (f 0)] + · rw [PMF.pure_apply, if_pos, mul_one, Fin.prod_univ_one] + funext i; rw [show i = 0 from Fin.ext (Nat.lt_one_iff.mp i.isLt)] + · intro b hb + rw [PMF.pure_apply, if_neg, mul_zero] + exact fun hc => hb ((congrFun hc 0).symm) + | succ n' IH => + intro s Δ ε₁ ε₂ l f + rw [privNoisedFamily, PMF.bind_apply, Fin.prod_univ_castSucc] + simp_rw [PMF.bind_apply, PMF.pure_apply] + conv_lhs => + enter [1, rest, 2] + rw [tsum_eq_single (f (Fin.last (n'+1))) fun b hb => by + rw [if_neg, mul_zero] + exact fun hc => hb ((fin_eq_iff_castSucc_last f rest b).mp hc).2.symm] + conv_lhs => enter [1, rest]; rw [show ∀ x y z : ENNReal, x * (y * z) = y * (x * z) from + fun _ _ _ => by ring] + rw [ENNReal.tsum_mul_left, + tsum_eq_single (fun j : Fin (n'+1) => f j.castSucc) fun b hb => by + rw [if_neg, mul_zero] + refine fun hc => hb (funext fun j => ?_) + exact (((fin_eq_iff_castSucc_last f b _).mp hc).1 j).symm] + have heq : f = (fun (i : Fin (n'+2)) => + if h : i.val < n'+1 then (fun (j : Fin (n'+1)) => f j.castSucc) ⟨i.val, h⟩ + else f (Fin.last (n'+1))) := + (fin_eq_iff_castSucc_last f _ _).mpr ⟨fun _ => rfl, rfl⟩ + rw [if_pos heq, mul_one, IH (fun i => s i.castSucc) Δ ε₁ ε₂ l (fun j => f j.castSucc)] + ring + +lemma tsum_pi_shift {n : ℕ} (Δfn : Fin (n+1) → ℤ) (g : (Fin (n+1) → ℤ) → ENNReal) : + (∑' f : Fin (n+1) → ℤ, g f) = (∑' f : Fin (n+1) → ℤ, g (fun i => f i + Δfn i)) := by + refine tsum_eq_tsum_of_ne_zero_bij (fun x i => x.1 i + Δfn i) ?_ ?_ (fun _ => rfl) + · intro x y hxy + funext i + have := congrFun hxy i + simp at this; omega + · intro x hx + refine ⟨⟨fun i => x i - Δfn i, ?_⟩, funext fun i => by simp⟩ + have heq : (fun i => x i - Δfn i + Δfn i) = (x : Fin (n+1) → ℤ) := funext fun i => by omega + simp only [Function.mem_support, heq]; exact hx + +def shiftToL₂ {n : ℕ} (iStar : Fin (n+1)) (s : Fin (n+1) → List T → ℤ) (Δ : ℕ+) + (l₁ l₂ : List T) : Fin (n+1) → ℤ := + fun i => if i = iStar then s iStar l₁ - s iStar l₂ - 2 * Δ else s i l₁ - s i l₂ + +@[simp] lemma shiftToL₂_self {n : ℕ} (iStar : Fin (n+1)) (s : Fin (n+1) → List T → ℤ) + (Δ : ℕ+) (l₁ l₂ : List T) : + shiftToL₂ iStar s Δ l₁ l₂ iStar = s iStar l₁ - s iStar l₂ - 2 * Δ := by simp [shiftToL₂] + +@[simp] lemma shiftToL₂_of_ne {n : ℕ} {iStar i : Fin (n+1)} (hi : i ≠ iStar) + (s : Fin (n+1) → List T → ℤ) (Δ : ℕ+) (l₁ l₂ : List T) : + shiftToL₂ iStar s Δ l₁ l₂ i = s i l₁ - s i l₂ := by simp [shiftToL₂, hi] + +lemma sensitivity_abs_le {q : List T → ℤ} {Δ : ℕ+} (Hsens : sensitivity q Δ) + {l₁ l₂ : List T} (Hn : Neighbour l₁ l₂) : |((q l₁ - q l₂ : ℤ))| ≤ (Δ : ℤ) := by + rw [Int.abs_eq_natAbs]; exact_mod_cast Hsens l₁ l₂ Hn + +lemma argmax_preserved_after_shift {n : ℕ} {s : Fin (n+1) → List T → ℤ} {Δ : ℕ+} + {l₁ l₂ : List T} (Hn : Neighbour l₁ l₂) (Hsens : ∀ i, sensitivity (s i) Δ) + (iStar : Fin (n+1)) {f : Fin (n+1) → ℤ} + (h : Fin.argmax n (fun i => f i + shiftToL₂ iStar s Δ l₁ l₂ i) = iStar) : + Fin.argmax n f = iStar := by + rw [Fin.argmax_eq_iff] at h ⊢ + obtain ⟨H_above, H_below⟩ := h + have abs_winner := abs_le.mp (sensitivity_abs_le (Hsens iStar) Hn) + refine ⟨fun j Hj => ?_, fun j Hj => ?_⟩ + all_goals + have hjne : j ≠ iStar := fun heq => by rw [heq] at Hj; exact lt_irrefl _ Hj + have abs_loser := abs_le.mp (sensitivity_abs_le (Hsens j) Hn) + · have := H_above j Hj + rw [shiftToL₂_of_ne hjne, shiftToL₂_self] at this + linarith + · have := H_below j Hj + rw [shiftToL₂_of_ne hjne, shiftToL₂_self] at this + linarith + +lemma loser_density_eq {q : List T → ℤ} {Δ ε₁ ε₂ : ℕ+} {l₁ l₂ : List T} (x : ℤ) : + (privNoisedQueryPure q (2*Δ) ε₁ ε₂ l₁) (x + (q l₁ - q l₂)) = + (privNoisedQueryPure q (2*Δ) ε₁ ε₂ l₂) x := by + simp only [privNoisedQueryPure, DiscreteLaplaceGenSamplePMF, DFunLike.coe] + rw [DiscreteLaplaceGenSample_periodic, DiscreteLaplaceGenSample_periodic] + congr 1; ring + +lemma winner_density_le {q : List T → ℤ} {Δ ε₁ ε₂ : ℕ+} {l₁ l₂ : List T} {ε : NNReal} + (HN : (ε₁ : NNReal) / ε₂ = ε) (x : ℤ) : + (privNoisedQueryPure q (2*Δ) ε₁ ε₂ l₁) (x + (q l₁ - q l₂ - 2 * Δ)) ≤ + ENNReal.ofReal (Real.exp ε) * (privNoisedQueryPure q (2*Δ) ε₁ ε₂ l₂) x := by + simp only [privNoisedQueryPure, DiscreteLaplaceGenSamplePMF, DFunLike.coe] + rw [DiscreteLaplaceGenSample_periodic, DiscreteLaplaceGenSample_periodic, + show (x + (↑(q l₁) - ↑(q l₂) - 2 * ↑Δ) - ↑(q l₁) : ℤ) = (x - ↑(q l₂)) + (-(2 * ↑Δ)) by ring] + have hlap := laplace_inequality_sub ε₁ ε₂ (x - ↑(q l₂)) (-(2 * ↑Δ)) (2 * Δ) + simp only [DiscreteLaplaceGenSamplePMF, DFunLike.coe, DiscreteLaplaceGenSample_periodic, + sub_zero] at hlap + refine hlap.trans (_root_.mul_le_mul_left ?_ _) + refine ENNReal.ofReal_le_ofReal (Real.exp_monotone (le_of_eq ?_)) + rw [← HN] + push_cast [abs_neg, abs_of_pos (by positivity : (0 : ℝ) < 2 * Δ)] + field_simp + +theorem privReportNoisyMax_DP (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (ε : NNReal) (HN : laplace_pureDP_noise_priv ε₁ ε₂ ε) + (Hsens : ∀ i, sensitivity (s i) Δ) : + PureDPSystem.prop (privReportNoisyMax n s Δ ε₁ ε₂) ε := by + unfold laplace_pureDP_noise_priv at HN + simp only [DPSystem.prop, PureDP] + apply singleton_to_event + unfold DP_singleton + intro l₁ l₂ Hn iStar + apply ENNReal.div_le_of_le_mul + unfold privReportNoisyMax + show (PMF.bind _ _) iStar ≤ _ * (PMF.bind _ _) iStar + rw [PMF.bind_apply, PMF.bind_apply] + simp_rw [PMF.pure_apply] + conv_lhs => enter [1, f]; rw [privNoisedFamily_apply n s Δ ε₁ ε₂ l₁ f] + conv_rhs => enter [2, 1, f]; rw [privNoisedFamily_apply n s Δ ε₁ ε₂ l₂ f] + rw [tsum_pi_shift (shiftToL₂ iStar s Δ l₁ l₂), ← ENNReal.tsum_mul_left] + refine ENNReal.tsum_le_tsum fun f => ?_ + by_cases hf : Fin.argmax n (fun i => f i + shiftToL₂ iStar s Δ l₁ l₂ i) = iStar + swap + · rw [if_neg fun h => hf h.symm, mul_zero]; exact _root_.zero_le _ + rw [if_pos hf.symm, if_pos (argmax_preserved_after_shift Hn Hsens iStar hf).symm, + mul_one, mul_one, + ← Finset.prod_erase_mul Finset.univ _ (Finset.mem_univ iStar), + ← Finset.prod_erase_mul Finset.univ + (fun i => (privNoisedQueryPure (s i) (2*Δ) ε₁ ε₂ l₂) (f i)) (Finset.mem_univ iStar)] + have h_losers : ∀ i ∈ Finset.univ.erase iStar, + (privNoisedQueryPure (s i) (2*Δ) ε₁ ε₂ l₁) (f i + shiftToL₂ iStar s Δ l₁ l₂ i) = + (privNoisedQueryPure (s i) (2*Δ) ε₁ ε₂ l₂) (f i) := fun i hi => by + rw [shiftToL₂_of_ne (Finset.ne_of_mem_erase hi)]; exact loser_density_eq _ + have h_winner : + (privNoisedQueryPure (s iStar) (2*Δ) ε₁ ε₂ l₁) (f iStar + shiftToL₂ iStar s Δ l₁ l₂ iStar) ≤ + ENNReal.ofReal (Real.exp ε) * (privNoisedQueryPure (s iStar) (2*Δ) ε₁ ε₂ l₂) (f iStar) := by + rw [shiftToL₂_self]; exact winner_density_le HN _ + rw [Finset.prod_congr rfl h_losers, mul_left_comm] + exact _root_.mul_le_mul_right h_winner _ + +theorem privReportNoisyMaxSPMF_DP (n : ℕ) (s : Fin (n+1) → List T → ℤ) (Δ ε₁ ε₂ : ℕ+) + (ε : NNReal) (HN : laplace_pureDP_noise_priv ε₁ ε₂ ε) + (Hsens : ∀ i, sensitivity (s i) Δ) : + PureDPSystem.prop (privReportNoisyMaxSPMF n s Δ ε₁ ε₂) ε := by + have heq : (fun l => privReportNoisyMaxSPMF n s Δ ε₁ ε₂ l) = + (fun l => (privReportNoisyMax n s Δ ε₁ ε₂ l : SPMF _)) := by + funext l + apply Subtype.ext + exact privReportNoisyMaxSLang_eq n s Δ ε₁ ε₂ l + rw [show privReportNoisyMaxSPMF n s Δ ε₁ ε₂ = _ from heq] + exact privReportNoisyMax_DP n s Δ ε₁ ε₂ ε HN Hsens + +end SLang + +end diff --git a/Test.lean b/Test.lean index cce67b1b..faf87e09 100644 --- a/Test.lean +++ b/Test.lean @@ -6,6 +6,8 @@ Authors: Jean-Baptiste Tristan import SampCert import SampCert.SLang import SampCert.Samplers.Gaussian.Properties +import SampCert.DifferentialPrivacy.Queries.MWEM.Basic +import SampCert.DifferentialPrivacy.Queries.MWEM.Updaters.Float import Init.Data.Float open SLang Std Int Array PMF @@ -269,7 +271,117 @@ def sparseVector_tests : IO Unit := do +def mwem_tests : IO Unit := do + let mwemBins : ℕ+ := 10 + let mwemQueries : Fin 10 → Fin mwemBins → ℤ := + fun i b => if i.val = b.val then 1 else 0 + have mwemBound : ∀ (i : Fin 10) (b : Fin mwemBins), |mwemQueries i b| ≤ ((1 : ℕ+) : ℤ) := by + intro i b + show |if i.val = b.val then (1 : ℤ) else 0| ≤ 1 + split <;> decide + let bias : List (Fin mwemBins) := + (List.replicate 200 (3 : Fin mwemBins)) ++ + (List.replicate 50 (1 : Fin mwemBins)) ++ + (List.replicate 50 (7 : Fin mwemBins)) + let nData := bias.length + let updater := floatMWUpdater (numBins := mwemBins) (n := 9) (Δ := 1) nData mwemQueries mwemBound + let ε₁ : ℕ+ := 1 + let ε₂ : ℕ+ := 1 + let T : ℕ := 5 + let A₀ : Fin mwemBins → Float := floatInit nData mwemBins + + IO.println s!"[mwem] testing float-valued MWEM, ({(ε₁ : ℕ)} / {(ε₂ : ℕ)})-DP per primitive, {T} rounds → {2 * T} * {(ε₁ : ℕ)}/{(ε₂ : ℕ)} total" + IO.println s!"data: {bias.length} records over {(mwemBins : ℕ)} bins, biased toward bin 3" + IO.println s!"workload: {10} indicator queries (one per bin)" + IO.println "" + + for trial in [:3] do + let transcript ← run <| mwemSPMF updater ε₁ ε₂ T A₀ bias + IO.println s!"#{trial} transcript: {transcript}" + IO.println "" + +def hardtBins : ℕ+ := 15 +def hardtNQ : ℕ := 12 +def hardtT : ℕ := 4 +def hardtN : ℕ := hardtNQ - 1 + +def hardtRange (i : Fin hardtNQ) : ℕ × ℕ := + let a : ℕ := i.val % 8 + let len : ℕ := (i.val / 8) * 4 + 3 + let c : ℕ := min (a + len) (hardtBins.val - 1) + (a, c) + +def hardtQueries : Fin hardtNQ → Fin hardtBins → ℤ := fun i b => + let (a, c) := hardtRange i + if a ≤ b.val ∧ b.val ≤ c then 1 else 0 + +theorem hardtBound : ∀ (i : Fin hardtNQ) (b : Fin hardtBins), + |hardtQueries i b| ≤ ((1 : ℕ+) : ℤ) := by + intro i b + show |if (hardtRange i).1 ≤ b.val ∧ b.val ≤ (hardtRange i).2 then (1 : ℤ) else 0| ≤ 1 + split <;> decide + +def hardtData : List (Fin hardtBins) := + let n := 200 + let make (i : ℕ) : Fin hardtBins := + ⟨ (i * 7 + 3) % 11 + 2, by simp [hardtBins]; omega ⟩ + (List.range n).map make + +def replayFloatUpdate (nData : ℕ) (updater : SyntheticUpdater (Fin hardtBins) hardtN 1 + (Fin hardtBins → Float) (floatInit nData hardtBins)) + (transcript : List (Fin hardtNQ × ℤ)) : Fin hardtBins → Float := + transcript.foldl (fun A (im : Fin hardtNQ × ℤ) => updater.update A im.1 im.2) + (floatInit nData hardtBins) + +def evalQueryOnFloat (q : Fin hardtBins → ℤ) (A : Fin hardtBins → Float) : Float := + ((List.finRange hardtBins).map (fun b => intToFloat (q b) * A b)).foldr (· + ·) 0.0 + +def evalQueryOnData (q : Fin hardtBins → ℤ) (D : List (Fin hardtBins)) : ℤ := + (D.map q).sum + +def hardt_test : IO Unit := do + let nData := hardtData.length + let updater := floatMWUpdater (numBins := hardtBins) (n := hardtN) (Δ := 1) nData hardtQueries hardtBound + let ε₁ : ℕ+ := 1 + let ε₂ : ℕ+ := 1 + let A₀ : Fin hardtBins → Float := floatInit nData hardtBins + + IO.println s!"[mwem-hardt] small-scale 1D range query experiment in the spirit of Hardt 2012, Section 3.1" + IO.println s!" domain: {(hardtBins : ℕ)} bins" + IO.println s!" workload: {hardtNQ} 1D range queries" + IO.println s!" data: {hardtData.length} synthetic records" + IO.println s!" T = {hardtT}, ε₁/ε₂ = {(ε₁ : ℕ)}/{(ε₂ : ℕ)} per primitive" + IO.println "" + + let nTrials := 3 + let mut total_avg_sq_err : Float := 0.0 + for trial in [:nTrials] do + IO.println s!" trial {trial}: running mwem..." + let transcript ← run <| mwemSPMF updater ε₁ ε₂ hardtT A₀ hardtData + IO.println s!" trial {trial}: transcript = {transcript}" + IO.println s!" trial {trial}: replaying updater..." + let A_T := replayFloatUpdate nData updater transcript + let A_T_norm := A_T + + IO.println s!" trial {trial}: evaluating workload..." + let mut sum_sq_err : Float := 0.0 + for i in (List.finRange hardtNQ) do + let q : Fin hardtBins → ℤ := hardtQueries i + let synth_ans := evalQueryOnFloat q A_T_norm + let true_ans : Float := intToFloat (evalQueryOnData q hardtData) + let err := synth_ans - true_ans + sum_sq_err := sum_sq_err + err * err + let avg_sq_err := sum_sq_err / intToFloat (hardtNQ : ℤ) + total_avg_sq_err := total_avg_sq_err + avg_sq_err + IO.println s!"#{trial} avg squared error per query: {avg_sq_err}" + + IO.println s!"mean over {nTrials} trials: {total_avg_sq_err / intToFloat (nTrials : ℤ)}" + IO.println s!"(units and parameter scale differ from Hardt's Figure 2; not a direct numerical comparison)" + IO.println "" + def main : IO Unit := do sparseVector_tests query_tests statistical_tests + mwem_tests + hardt_test