Commit 7cb475a9 authored by dsbowen's avatar dsbowen
Browse files

Fixed bug in empirical Bayes prior distribution plot in Bayes primer notebook

parent 68a258cc
Loading
Loading
Loading
Loading
Loading
+277 −80

File changed.

Preview size limit exceeded, changes collapsed.

+21 −0
Original line number Diff line number Diff line
@@ -8,6 +8,7 @@ from typing import Any, Dict, Tuple, Union

import numpy as np
from scipy.optimize import minimize_scalar
from scipy.stats import multivariate_normal

from ..base import ColumnsType, Numeric1DArray
from .base import BayesModelBase, BayesResults
@@ -328,6 +329,26 @@ class LinearEmpiricalBayes(EmpiricalBayesBase):

        return prior_mean_params, prior_cov_params

    def prior_mean_rvs(self, size: int = 1) -> np.ndarray:
        """Sample from the distribution of prior means.

        Args:
            size (int, optional): Number of samples to draw. Defaults to 1.

        Returns:
            np.ndarray: (size, n) array of prior mean samples.
        """
        # TODO: incorporate estimate_prior_params keyword arguments
        # possibly pass in a prior_cov parameter to be consistent with heirarchical Bayes
        _, prior_cov_params = self.estimate_prior_params()
        prior_cov = self.estimate_prior_cov(prior_cov_params)
        X_T = self.X.T
        tau_inv = np.linalg.inv(prior_cov + self.cov)
        XT_tauinv_X_inv = np.linalg.inv(X_T @ tau_inv @ self.X)
        beta_bar = XT_tauinv_X_inv @ X_T @ tau_inv @ self.mean
        beta = multivariate_normal.rvs(beta_bar, XT_tauinv_X_inv, size=size)
        return (self.X @ beta.reshape(1, -1)).squeeze()

    def _estimate_prior_mean_params(self, prior_cov: np.ndarray) -> np.ndarray:
        """Estimate prior mean parameter vector.

+4 −4
Original line number Diff line number Diff line
@@ -236,10 +236,10 @@ class LinearHierarchicalBayes(HierarchicalBayesBase):
        """
        X_T = self.X.T
        tau_inv = np.linalg.inv(prior_cov + self.cov)
        XT_tauinv_X = X_T @ tau_inv @ self.X
        beta_bar = np.linalg.inv(XT_tauinv_X) @ X_T @ tau_inv @ self.mean
        beta = multivariate_normal.rvs(beta_bar, np.linalg.inv(XT_tauinv_X), size=size)
        return (self.X @ beta.reshape(-1, 1)).squeeze()
        XT_tauinv_X_inv = np.linalg.inv(X_T @ tau_inv @ self.X)
        beta_bar = XT_tauinv_X_inv @ X_T @ tau_inv @ self.mean
        beta = multivariate_normal.rvs(beta_bar, XT_tauinv_X_inv, size=size)
        return (self.X @ beta.reshape(1, -1)).squeeze()

    def _scaled_log_likelihood(self, prior_cov: np.ndarray) -> float:
        # compute the scaled log likelihood; see HierarchicalBayesBase
+7 −0
Original line number Diff line number Diff line
@@ -60,3 +60,10 @@ class TestResults:

    def test_reconstruction_point_plot(self, results):
        results.reconstruction_point_plot()


def test_prior_mean_rvs(size=10):
    # TODO: test with estimate_prior_params keyword arguments
    model = LinearEmpiricalBayes(mean, cov)
    assert model.prior_mean_rvs().shape == (n_policies,)
    assert model.prior_mean_rvs(size).shape == (n_policies, size)