log_lik <- function(x, mu, w) sum(dnorm(x, mu, exp(w / 2), log = TRUE))
log_prior <- function(mu, w, mu_0, sigma_mu, w_0, sigma_w) {
dnorm(mu, mu_0, sigma_mu, log = TRUE) + dnorm(w, w_0, sigma_w, log = TRUE)
}
neg_log_post <- function(theta, x, mu_0, sigma_mu, w_0, sigma_w) {
-log_lik(x, theta[1], theta[2]) -
log_prior(theta[1], theta[2], mu_0, sigma_mu, w_0, sigma_w)
}
log_dmvnorm_batch <- function(Theta, mu, A) { # log N(mu, A^-1) at every column at once
Tc <- Theta - mu
quad <- colSums((A %*% Tc) * Tc)
0.5 * (-nrow(Theta) * log(2 * pi) + sum(log(svd(A)$d)) - quad)
}
log_marlik_imps <- function(x, mu_0, sigma_mu, w_0, sigma_w, S) {
fit <- nlm(neg_log_post, c(mean(x), log(stats::var(x))), hessian = TRUE,
x = x, mu_0 = mu_0, sigma_mu = sigma_mu, w_0 = w_0, sigma_w = sigma_w)
mu_hat <- fit$estimate; H <- fit$hessian
L <- t(chol(solve(H)))
Theta <- L %*% matrix(rnorm(2 * S), 2, S) + mu_hat
log_g <- log_dmvnorm_batch(Theta, mu_hat, H)
log_f <- -apply(Theta, 2, neg_log_post, x = x, mu_0 = mu_0, sigma_mu = sigma_mu,
w_0 = w_0, sigma_w = sigma_w)
lw <- log_f - log_g
m <- max(lw); v <- exp(lw - m)
structure(log_sum_exp(lw) - log(S), ess = log_ess(lw), se = sd(v) / (sqrt(S) * mean(v)))
}
log_marlik_prior <- function(x, mu_0, sigma_mu, w_0, sigma_w, S) {
mus <- rnorm(S, mu_0, sigma_mu); ws <- rnorm(S, w_0, sigma_w)
lw <- vapply(seq_len(S), function(i) log_lik(x, mus[i], ws[i]), numeric(1))
m <- max(lw); v <- exp(lw - m)
structure(log_sum_exp(lw) - log(S), ess = log_ess(lw), se = sd(v) / (sqrt(S) * mean(v)))
}
log_marlik_grid <- function(x, mu_0, sigma_mu, w_0, sigma_w, n_grid = 400, width = 8) {
n <- length(x)
fit <- nlm(neg_log_post, c(mean(x), log(stats::var(x))), hessian = TRUE,
x = x, mu_0 = mu_0, sigma_mu = sigma_mu, w_0 = w_0, sigma_w = sigma_w)
se <- sqrt(diag(solve(fit$hessian)))
mu_grid <- fit$estimate[1] + seq(-width, width, length.out = n_grid) * se[1]
w_grid <- fit$estimate[2] + seq(-width, width, length.out = n_grid) * se[2]
SS <- sum(x^2) - 2 * mu_grid * sum(x) + n * mu_grid^2
log_lik_grid <- outer(SS, w_grid,
function(ss, w) -0.5 * n * log(2 * pi) - n * w / 2 - 0.5 * exp(-w) * ss)
log_prior_grid <- outer(dnorm(mu_grid, mu_0, sigma_mu, log = TRUE),
dnorm(w_grid, w_0, sigma_w, log = TRUE), "+")
h_mu <- diff(mu_grid[1:2]); h_w <- diff(w_grid[1:2])
log_sum_exp(as.vector(log_lik_grid + log_prior_grid)) + log(h_mu) + log(h_w)
}