Skip to content
Open
Show file tree
Hide file tree
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
6 changes: 6 additions & 0 deletions CHANGELOG_UNRELEASED.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
136 changes: 136 additions & 0 deletions theories/probability_theory/normal_distribution.v
Original file line number Diff line number Diff line change
Expand Up @@ -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 *)
(* ``` *)
(* *)
(******************************************************************************)
Expand Down Expand Up @@ -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 ->
Expand Down Expand Up @@ -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.
Loading