From 4cf93f85dcd16e88869e28dda97ca45c8012684b Mon Sep 17 00:00:00 2001 From: Guillaume Baudart Date: Mon, 27 Apr 2026 10:58:54 +0900 Subject: [PATCH] Gaussian-Gaussian conjugate prior (#43) * add Gaussian-Gaussian conjugate prior Developed with the assistance of Claude Opus 4.8 and Fable 5 (1M context). --------- Co-authored-by: Reynald Affeldt --- CHANGELOG_UNRELEASED.md | 6 + .../probability_theory/normal_distribution.v | 136 ++++++++++++++++++ 2 files changed, 142 insertions(+) diff --git a/CHANGELOG_UNRELEASED.md b/CHANGELOG_UNRELEASED.md index 67bb43c3b6..496e58a254 100644 --- a/CHANGELOG_UNRELEASED.md +++ b/CHANGELOG_UNRELEASED.md @@ -4,6 +4,12 @@ ### Added +- in `normal_distribution.v`: + + definition `stddev_post` + + lemma `stddev_post_neq0` + + definition `mean_post` + + lemmas `normal_fun_conjugate`, `normal_pdf_conjugate`, `normal_prob_conjugate` + ### Changed ### Renamed diff --git a/theories/probability_theory/normal_distribution.v b/theories/probability_theory/normal_distribution.v index f336e75190..d4caec1b3b 100644 --- a/theories/probability_theory/normal_distribution.v +++ b/theories/probability_theory/normal_distribution.v @@ -18,6 +18,8 @@ From mathcomp Require Import lebesgue_integral ftc gauss_integral. (* standard deviation s *) (* Using normal_peak and normal_pdf. *) (* normal_prob m s == normal probability measure *) +(* stddev_post s0 s == posterior standard deviation *) +(* mean_post m s0 x s == posterior mean given an observation x *) (* ``` *) (* *) (******************************************************************************) @@ -596,6 +598,7 @@ Section normal_prob_lemmas. Context {R : realType}. Local Notation mu := lebesgue_measure. Local Open Scope ereal_scope. + Import MeasurableR. Lemma integral_normal_prob (m s : R) f U : measurable U -> @@ -833,3 +836,136 @@ by rewrite sqrtrM ?addr_ge0 ?sqr_ge0// mulrC. Qed. End normal_probD. + +Section gaussian_conjugate. +Context {R : realType}. +Implicit Types (sigma x theta : R) (V : set R). + +(**md $\sqrt{\frac{\sigma_0^2 \sigma^2}{\sigma_0^2 + \sigma^2}}$ *) +Definition stddev_post sigma0 sigma : R := + Num.sqrt (sigma0 ^+ 2 * sigma ^+ 2 / (sigma0 ^+ 2 + sigma ^+ 2)). + +Lemma stddev_post_neq0 sigma0 sigma : + sigma0 != 0 -> sigma != 0 -> stddev_post sigma0 sigma != 0. +Proof. +move=> sigma0_neq0 sigma_neq0. +rewrite lt0r_neq0 ?sqrtr_gt0 ?divr_gt0 ?addr_gt0 ?exprn_even_gt0//. +by rewrite mulr_gt0// exprn_even_gt0. +Qed. + +(**md $\frac{\sigma^2 \mu_0 + \sigma_0^2 x}{\sigma_0^2 + \sigma^2}$ *) +Definition mean_post (mu0 : R) sigma0 x sigma : R := + (sigma ^+ 2 * mu0 + sigma0 ^+ 2 * x) / (sigma0 ^+ 2 + sigma ^+ 2). + +(* "complete the square" for the normal_fun exponents *) +Lemma normal_fun_conjugate (mu0 : R) sigma0 sigma x theta : + sigma0 != 0 -> sigma != 0 -> + normal_fun theta sigma x * normal_fun mu0 sigma0 theta = + normal_fun mu0 (Num.sqrt (sigma0 ^+ 2 + sigma ^+ 2)) x * + normal_fun (mean_post mu0 sigma0 x sigma) (stddev_post sigma0 sigma) theta. +Proof. +move=> sigma0_neq0 sigma_neq0. +have ? : 0 < sigma0 ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma0 ^+ 2 + sigma ^+ 2 by rewrite addr_gt0. +rewrite /normal_fun -2!expRD; congr (expR _). +rewrite (sqr_sqrtr (ltW _))// /stddev_post sqr_sqrtr. + by rewrite (divr_ge0 _ (ltW _))// mulr_ge0// ltW. +by rewrite /mean_post; field by rewrite ?sigma0_neq0 ?sigma_neq0 ?gt_eqF. +Qed. + +(**md "complete the square": $p(x | \theta) * p(\theta)$ factors as a theta-free + marginal times the posterior density *) +Lemma normal_pdf_conjugate mu0 sigma0 sigma x theta : + sigma0 != 0 -> sigma != 0 -> + normal_pdf theta sigma x * normal_pdf mu0 sigma0 theta = + normal_pdf mu0 (Num.sqrt (sigma0 ^+ 2 + sigma ^+ 2)) x * + normal_pdf (mean_post mu0 sigma0 x sigma) (stddev_post sigma0 sigma) theta. +Proof. +move=> sigma0_neq0 sigma_neq0. +have ? : 0 < sigma0 ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma0 ^+ 2 + sigma ^+ 2 by apply: addr_gt0. +have ? : Num.sqrt (sigma0 ^+ 2 + sigma ^+ 2) != 0 by rewrite sqrtr_eq0 -ltNge. +have ? : 0 < stddev_post sigma0 sigma. + by rewrite /stddev_post sqrtr_gt0 divr_gt0// mulr_gt0. +have ? := stddev_post_neq0 sigma0_neq0 sigma_neq0. +rewrite !normal_pdfE // mulrACA [RHS]mulrACA. +congr (_ * _); first last. + exact: normal_fun_conjugate. +have ? : 0 <= sigma ^+ 2 * pi *+ 2 + by rewrite mulrn_wge0// mulr_ge0 ?pi_ge0// ltW. +have ? : 0 <= sigma0 ^+ 2 * pi *+ 2 + by rewrite mulrn_wge0// mulr_ge0 ?pi_ge0// ltW. +have ? : 0 <= (sigma0 ^+ 2 + sigma ^+ 2) * pi *+ 2 + by rewrite mulrn_wge0// mulr_ge0 ?pi_ge0// ltW. +rewrite /normal_peak (_ : Num.sqrt _ ^+ 2 = sigma0 ^+ 2 + sigma ^+ 2). + by rewrite sqr_sqrtr// ltW. +rewrite (_ : stddev_post sigma0 sigma ^+ 2 + = sigma0 ^+ 2 * sigma ^+ 2 / (sigma0 ^+ 2 + sigma ^+ 2)). + by rewrite /stddev_post sqr_sqrtr// divr_ge0// ?addr_ge0 ?ltW// mulr_gt0. +rewrite -invfM -[RHS]invfM -!sqrtrM//. +by congr (Num.sqrt _)^-1; field by rewrite gt_eqF. +Qed. + +Import MeasurableR. + +(* bounded integrand against a probability measure *) +Local Lemma integrable_normal_pdf_likelihood mu0 sigma0 sigma x V : + sigma != 0 -> measurable V -> + (normal_prob mu0 sigma0).-integrable V + (fun theta => (normal_pdf theta sigma x)%:E). +Proof. +move=> sigma_neq0 mV. +have pdf_sym theta : normal_pdf theta sigma x = normal_pdf x sigma theta. + rewrite !normal_pdfE //; congr (_ * _). + by rewrite /normal_fun -[in LHS](opprB theta x) sqrrN. +have -> : (fun theta => (normal_pdf theta sigma x)%:E) = + EFin \o normal_pdf x sigma by apply/funext => theta; rewrite pdf_sym. +apply: (@measurable_bounded_integrable _ _ _ + (normal_prob mu0 sigma0) + (normal_pdf x sigma) _ mV). +- by apply: (le_lt_trans (probability_le1 _ _)) => //; rewrite ltey. +- by apply: measurable_funTS; exact: measurable_normal_pdf. +- exists (normal_peak sigma); split; first by rewrite num_real. + move=> y ynp /= theta _. + rewrite ger0_norm; first exact: normal_pdf_ge0. + exact: le_trans (normal_pdf_ub _ _ sigma_neq0) (ltW ynp). +Qed. + +(* Gaussian-Gaussian conjugate prior: the posterior of theta given X = x is + the single Gaussian normal_prob (mean_post ..) (stddev_post ..) *) +Lemma normal_prob_conjugate mu0 sigma0 sigma x V : + sigma0 != 0 -> sigma != 0 -> measurable V -> + ((\int[normal_prob mu0 sigma0]_(theta in V) (normal_pdf theta sigma x)%:E) / + \int[normal_prob mu0 sigma0]_theta (normal_pdf theta sigma x)%:E = + normal_prob (mean_post mu0 sigma0 x sigma) (stddev_post sigma0 sigma) V)%E. +Proof. +move=> sigma0_neq0 sigma_neq0 mV. +have ? : 0 < sigma0 ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma ^+ 2 by rewrite exprn_even_gt0. +have ? : 0 < sigma0 ^+ 2 + sigma ^+ 2 by apply: addr_gt0. +have ? : Num.sqrt (sigma0 ^+ 2 + sigma ^+ 2) != 0 by rewrite sqrtr_eq0 -ltNge. +pose K := normal_pdf mu0 (Num.sqrt (sigma0 ^+ 2 + sigma ^+ 2)) x. +have Kpos : 0 < K. + by rewrite /K normal_pdfE// mulr_gt0 ?normal_peak_gt0//; exact: expR_gt0. +have ? : (K%:E != 0)%E by rewrite eqe gt_eqF. +have step U : measurable U -> + (\int[normal_prob mu0 sigma0]_(theta in U) (normal_pdf theta sigma x)%:E = + K%:E * + normal_prob (mean_post mu0 sigma0 x sigma) (stddev_post sigma0 sigma) U)%E. + move=> mU. + rewrite integral_normal_prob ?integrable_normal_pdf_likelihood//=. + under eq_integral => theta _ do + rewrite -EFinM + (normal_pdf_conjugate mu0 x theta sigma0_neq0 sigma_neq0) + EFinM. + rewrite -/K ge0_integralZl_EFin //=. + - by move=> theta _; rewrite lee_fin; exact: normal_pdf_ge0. + - by apply/measurable_EFinP/measurable_funTS; exact: measurable_normal_pdf. + - exact: ltW. +rewrite (step _ mV) (step _ measurableT) probability_setT mule1. +by rewrite muleAC divee // mul1e. +Qed. + +End gaussian_conjugate.