R/diagnostic_posterior.R

Defines functions diagnostic_posterior.blavaan diagnostic_posterior.CmdStanFit diagnostic_posterior.stanfit diagnostic_posterior.brmsfit diagnostic_posterior.stanmvreg diagnostic_posterior.stanreg diagnostic_posterior.default diagnostic_posterior

Documented in diagnostic_posterior diagnostic_posterior.default diagnostic_posterior.stanreg

#' Posteriors Sampling Diagnostic
#'
#' Extract diagnostic metrics (Effective Sample Size (`ESS`), `Rhat` and Monte
#' Carlo Standard Error `MCSE`).
#'
#' @param posterior A `stanreg`, `stanfit`, `brmsfit`, or `blavaan` object; a
#' list of data frames or matrices representing MCMC chains (rows as samples,
#' columns as parameters); or a 3D array (dimensions: samples, chains,
#' parameters)
#' @param diagnostic Diagnostic metrics to compute.  Character (vector) or list
#' with one or more of these options: `"ESS"`, `"ESS_bulk"`, `"Rhat"`, `"MCSE"`
#' or `"all"`. `"ESS"` returns the **tail-ESS** (the minimum of the effective
#' sample sizes for the 5% and 95% quantiles), which is the most relevant
#' diagnostic for assessing the reliability of credible intervals and other
#' tail-based quantities. `"ESS_bulk"` additionally returns the **bulk-ESS**
#' (the effective sample size for the bulk of the posterior, useful for
#' assessing the reliability of central tendency estimates such as the mean or
#' median). `"all"` includes both tail and bulk `"ESS"`, `"Rhat"`, and `"MCSE"`.
#'
#' @inheritSection hdi Model components
#'
#' @details
#'   **Effective Sample (ESS)** should be as large as possible, although for
#'   most applications, an effective sample size greater than 1000 is sufficient
#'   for stable estimates (_Bürkner, 2017_). The ESS returned by
#'   `diagnostic_posterior()` is the **tail-ESS**: it corresponds to the
#'   minimum of the effective sample sizes for the 5% and 95% quantiles, and
#'   is a diagnostic for the sampling efficiency in the tails of the posterior
#'   distribution. It is more relevant than the bulk-ESS for assessing
#'   the reliability of credible intervals, probabilities of direction, and
#'   other tail-based quantities. Note that the tail-ESS may differ from the
#'   ESS reported by `brms` (`Bulk_ESS`) or other tools; use `"ESS_bulk"` to
#'   also retrieve the bulk-ESS.
#'
#'   **Rhat** should be the closest to 1. It should not be larger than 1.1
#'   (_Gelman and Rubin, 1992_) or 1.01 (_Vehtari et al., 2019_). The split
#'   Rhat statistic quantifies the consistency of an ensemble of Markov chains.
#'
#'   **Monte Carlo Standard Error (MCSE)** is another measure of accuracy of the
#'   chains. It is defined as standard deviation of the chains divided by their
#'   effective sample size (the formula for `mcse()` is from Kruschke 2015, p.
#'   187). The MCSE "provides a quantitative suggestion of how big the estimation
#'   noise is".
#'
#'
#' @examplesIf require("rstanarm") && require("brms")
#' \donttest{
#' # rstanarm models
#' # -----------------------------------------------
#' model <- suppressWarnings(
#'   rstanarm::stan_glm(mpg ~ wt + gear, data = mtcars, chains = 2, iter = 200, refresh = 0)
#' )
#' diagnostic_posterior(model)
#'
#' # brms models
#' # -----------------------------------------------
#' model <- brms::brm(mpg ~ wt + cyl, data = mtcars)
#' diagnostic_posterior(model)
#' }
#' @examplesIf require("rstan")
#' set.seed(101)
#' mkdata <- function(nrow = 1000, ncol = 2, parnm = LETTERS[1:ncol]) {
#'   x <- as.data.frame(replicate(ncol, rnorm(nrow)))
#'   names(x) <- parnm
#'   x
#' }
#' dd <- replicate(5, mkdata(), simplify = FALSE)
#' diagnostic_posterior(dd)
#' @references
#' - Gelman, A., & Rubin, D. B. (1992). Inference from iterative simulation
#'   using multiple sequences. Statistical science, 7(4), 457-472.
#' - Vehtari, A., Gelman, A., Simpson, D., Carpenter, B., and Bürkner, P. C.
#'   (2019). Rank-normalization, folding, and localization: An improved Rhat
#'   for assessing convergence of MCMC. arXiv preprint arXiv:1903.08008.
#' - Kruschke, J. (2014). Doing Bayesian data analysis: A tutorial with R,
#'   JAGS, and Stan. Academic Press.
#' @export
diagnostic_posterior <- function(posterior, ...) {
  UseMethod("diagnostic_posterior")
}


#' @rdname diagnostic_posterior
#' @export
diagnostic_posterior.default <- function(posterior, diagnostic = "all", ...) {
  ## check input, coerce to array
  if (is.list(posterior)) {
    for (i in seq_along(posterior)) {
      p <- posterior[[i]]
      if (
        !((inherits(p, "data.frame") || inherits(p, "matrix")) &&
          (length(dim(p)) == 2) &&
          (ncol(p) == ncol(posterior[[1]])))
      ) {
        insight::format_error(
          "'posterior' must be a 3D array or a list of data frames with equal numbers of columns."
        )
      }
    }
    insight::check_if_installed("posterior")
    posterior <- posterior::as_draws_array(posterior)
    ## draws_array() class messes things up downstream ...
    class(posterior) <- "array"
  }
  if (!(inherits(posterior, "array") && length(dim(posterior)) == 3)) {
    insight::format_error("Expecting a 3D array for 'posterior'.")
  }

  ret <- data.frame(Parameter = colnames(posterior[[1]]))
  if (is.null(diagnostic)) {
    return(ret)
  }

  .diag_opts <- c("Rhat", "ESS", "ESS_bulk", "MCSE")
  if (diagnostic == "all") {
    diagnostic <- c("Rhat", "ESS", "MCSE")
  }

  ## need ESS for MCSE, so compute these in any case
  insight::check_if_installed("rstan")
  mon <- rstan::monitor(posterior, print = FALSE, probs = 0.5)
  mon_df <- as.data.frame(mon)
  # Use Tail_ESS if available (rstan >= 2.21), otherwise fall back to n_eff
  ess_col <- if ("Tail_ESS" %in% names(mon_df)) {
    round(mon_df[["Tail_ESS"]])
  } else {
    mon_df[["n_eff"]]
  }
  ret <- data.frame(
    Parameter = rownames(mon),
    ESS = ess_col,
    Rhat = mon_df[["Rhat"]],
    MCSE = mon_df[["MCSE_Q50"]]
  )
  if ("ESS_bulk" %in% diagnostic && "Bulk_ESS" %in% names(mon_df)) {
    ret$ESS_bulk <- round(mon_df[["Bulk_ESS"]])
  }
  ret[c("Parameter", intersect(diagnostic, names(ret)))]
}

#' @inheritParams insight::get_parameters.BFBayesFactor
#' @inheritParams insight::get_parameters
#' @rdname diagnostic_posterior
#' @export
diagnostic_posterior.stanreg <- function(
  posterior,
  diagnostic = "all",
  effects = "fixed",
  component = "location",
  parameters = NULL,
  ...
) {
  # Find parameters
  params <- insight::find_parameters(
    posterior,
    effects = effects,
    component = component,
    parameters = parameters,
    flatten = TRUE
  )

  # If no diagnostic
  if (is.null(diagnostic)) {
    return(data.frame(Parameter = params))
  }

  diagnostic <- match.arg(
    diagnostic,
    c("ESS", "ESS_bulk", "Rhat", "MCSE", "all"),
    several.ok = TRUE
  )

  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS_tail", "ESS_bulk", "Rhat", "MCSE", "khat")
  } else {
    diagnostic <- diagnostic
    if ("Rhat" %in% diagnostic) {
      diagnostic <- c(diagnostic, "khat")
    }
  }
  # ESS: use tail ESS by default, with optional bulk ESS
  if ("ESS" %in% diagnostic) {
    diagnostic[diagnostic == "ESS"] <- "ESS_tail"
  }

  # Get indices and rename
  diagnostic_df <- as.data.frame(posterior$stan_summary)
  diagnostic_df$Parameter <- row.names(diagnostic_df)
  # special handling for MCSE, due to some parameters (like lp__) missing in rows
  MCSE <- mcse(posterior, effects = "full")
  diagnostic_df <- merge(diagnostic_df, MCSE, by = "Parameter", all = FALSE)

  # ESS: use tail ESS by default, with optional bulk ESS
  if (any(c("ESS_tail", "ESS_bulk") %in% diagnostic)) {
    ess_data <- effective_sample(
      posterior,
      effects = effects,
      component = component,
      parameters = parameters
    )
    if ("ESS_tail" %in% diagnostic && "ESS_tail" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_tail <- stats::setNames(ess_data$ESS_tail, ess_data$Parameter)[diagnostic_df$Parameter]
    }
    if ("ESS_bulk" %in% diagnostic && "ESS_bulk" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_bulk <- stats::setNames(ess_data$ESS_bulk, ess_data$Parameter)[diagnostic_df$Parameter]
    }
  }

  # Select columns
  available_columns <- intersect(colnames(diagnostic_df), c("Parameter", diagnostic))
  diagnostic_df <- diagnostic_df[available_columns]
  names(diagnostic_df)[available_columns == "khat"] <- "Khat"
  row.names(diagnostic_df) <- NULL

  # Remove columns with all Nans
  diagnostic_df <- diagnostic_df[!sapply(diagnostic_df, function(x) all(is.na(x)))]

  # Select rows
  diagnostic_df[diagnostic_df$Parameter %in% params, ]
}


#' @inheritParams insight::get_parameters
#' @export
diagnostic_posterior.stanmvreg <- function(
  posterior,
  diagnostic = "all",
  effects = "fixed",
  parameters = NULL,
  ...
) {
  # Find parameters
  all_params <- insight::find_parameters(
    posterior,
    effects = effects,
    parameters = parameters,
    flatten = FALSE
  )

  params <- unlist(
    lapply(names(all_params), function(i) {
      all_params[[i]]$sigma <- NULL
      unlist(all_params[[i]], use.names = FALSE)
    }),
    use.names = FALSE
  )

  # If no diagnostic
  if (is.null(diagnostic)) {
    return(data.frame(Parameter = params))
  }

  diagnostic <- match.arg(
    diagnostic,
    c("ESS", "ESS_bulk", "Rhat", "MCSE", "all"),
    several.ok = TRUE
  )

  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS_tail", "ESS_bulk", "Rhat", "MCSE", "khat")
  } else {
    diagnostic <- diagnostic
    if ("Rhat" %in% diagnostic) diagnostic <- c(diagnostic, "khat")
  }
  # ESS: use tail ESS by default, with optional bulk ESS
  if ("ESS" %in% diagnostic) {
    diagnostic[diagnostic == "ESS"] <- "ESS_tail"
  }

  # Get indices and rename
  diagnostic_df <- as.data.frame(posterior$stan_summary)
  diagnostic_df$Parameter <- row.names(diagnostic_df)
  # special handling for MCSE, due to some parameters (like lp__) missing in rows
  MCSE <- mcse(posterior, effects = effects)
  diagnostic_df <- merge(diagnostic_df, MCSE, by = "Parameter", all = FALSE)

  # ESS: use tail ESS by default, with optional bulk ESS
  if (any(c("ESS_tail", "ESS_bulk") %in% diagnostic)) {
    ess_data <- effective_sample(
      posterior,
      effects = effects,
      parameters = parameters
    )
    if ("ESS_tail" %in% diagnostic && "ESS_tail" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_tail <- stats::setNames(ess_data$ESS_tail, ess_data$Parameter)[diagnostic_df$Parameter]
    }
    if ("ESS_bulk" %in% diagnostic && "ESS_bulk" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_bulk <- stats::setNames(ess_data$ESS_bulk, ess_data$Parameter)[diagnostic_df$Parameter]
    }
  }

  # Select columns
  available_columns <- intersect(colnames(diagnostic_df), c("Parameter", diagnostic))
  diagnostic_df <- diagnostic_df[available_columns]
  names(diagnostic_df)[available_columns == "khat"] <- "Khat"
  row.names(diagnostic_df) <- NULL

  # Remove columns with all Nans
  diagnostic_df <- diagnostic_df[!sapply(diagnostic_df, function(x) all(is.na(x)))]

  diagnostic_df$Response <- gsub("(b\\[)*(.*)\\|(.*)", "\\2", diagnostic_df$Parameter)

  # Select rows
  diagnostic_df <- diagnostic_df[diagnostic_df$Parameter %in% params, ]

  # clean parameters
  for (i in unique(diagnostic_df$Response)) {
    diagnostic_df$Parameter <- gsub(
      sprintf("%s|", i),
      "",
      diagnostic_df$Parameter,
      fixed = TRUE
    )
  }

  diagnostic_df
}


#' @inheritParams insight::get_parameters
#' @export
diagnostic_posterior.brmsfit <- function(
  posterior,
  diagnostic = "all",
  effects = "fixed",
  component = "conditional",
  parameters = NULL,
  ...
) {
  # Find parameters
  params <- insight::find_parameters(
    posterior,
    effects = effects,
    component = component,
    parameters = parameters,
    flatten = TRUE
  )

  # If no diagnostic
  if (is.null(diagnostic)) {
    return(data.frame(Parameter = params))
  }

  # Get diagnostic
  diagnostic <- match.arg(
    diagnostic,
    c("ESS", "ESS_bulk", "Rhat", "MCSE", "all"),
    several.ok = TRUE
  )

  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS_tail", "ESS_bulk", "Rhat", "MCSE")
  }
  # ESS: use tail ESS by default, with optional bulk ESS
  if ("ESS" %in% diagnostic) {
    diagnostic[diagnostic == "ESS"] <- "ESS_tail"
  }

  # Initialize diagnostic dataframe
  diagnostic_df <- data.frame(Parameter = params, stringsAsFactors = FALSE)

  # Use posterior::summarise_draws() as single source for Rhat and ESS
  if (any(c("ESS_tail", "ESS_bulk", "Rhat") %in% diagnostic)) {
    insight::check_if_installed("posterior")
    idx <- as.data.frame(posterior::summarise_draws(posterior))
    idx <- idx[idx$variable %in% params, ]

    if ("Rhat" %in% diagnostic) {
      rhat_df <- data.frame(
        Parameter = idx$variable,
        Rhat = idx$rhat,
        stringsAsFactors = FALSE
      )
      diagnostic_df <- merge(diagnostic_df, rhat_df, by = "Parameter", all.x = TRUE)
    }

    if ("ESS_tail" %in% diagnostic) {
      ess_df <- data.frame(
        Parameter = idx$variable,
        ESS_tail = round(idx$ess_tail),
        stringsAsFactors = FALSE
      )
      diagnostic_df <- merge(diagnostic_df, ess_df, by = "Parameter", all.x = TRUE)
    }

    if ("ESS_bulk" %in% diagnostic) {
      ess_bulk_df <- data.frame(
        Parameter = idx$variable,
        ESS_bulk = round(idx$ess_bulk),
        stringsAsFactors = FALSE
      )
      diagnostic_df <- merge(diagnostic_df, ess_bulk_df, by = "Parameter", all.x = TRUE)
    }
  }

  # MCSE
  if ("MCSE" %in% diagnostic) {
    MCSE <- mcse(
      posterior,
      effects = effects,
      component = component,
      parameters = parameters
    )
    diagnostic_df <- merge(diagnostic_df, MCSE, by = "Parameter", all.x = TRUE)
  }

  row.names(diagnostic_df) <- NULL

  # Remove columns with all NAs
  diagnostic_df <- diagnostic_df[!sapply(diagnostic_df, function(x) all(is.na(x)))]

  # Select rows
  diagnostic_df[diagnostic_df$Parameter %in% params, ]
}


#' @inheritParams insight::get_parameters
#' @export
diagnostic_posterior.stanfit <- function(
  posterior,
  diagnostic = "all",
  effects = "fixed",
  parameters = NULL,
  ...
) {
  # Find parameters
  params <- insight::find_parameters(
    posterior,
    effects = effects,
    parameters = parameters,
    flatten = TRUE
  )

  # If no diagnostic
  if (is.null(diagnostic)) {
    return(data.frame(Parameter = params))
  }

  # Get diagnostic
  diagnostic <- match.arg(
    diagnostic,
    c("ESS", "ESS_bulk", "Rhat", "MCSE", "all"),
    several.ok = TRUE
  )
  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS_tail", "ESS_bulk", "Rhat", "MCSE")
  }
  # ESS: use tail ESS by default, with optional bulk ESS
  if ("ESS" %in% diagnostic) {
    diagnostic[diagnostic == "ESS"] <- "ESS_tail"
  }

  insight::check_if_installed("rstan")

  all_params <- insight::find_parameters(posterior, effects = effects, flatten = TRUE)

  diagnostic_df <- data.frame(
    Parameter = all_params,
    stringsAsFactors = FALSE
  )

  if (any(c("ESS_tail", "ESS_bulk") %in% diagnostic)) {
    ess_data <- effective_sample(posterior, effects = effects, parameters = parameters)
    if ("ESS_tail" %in% diagnostic && "ESS_tail" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_tail <- stats::setNames(ess_data$ESS_tail, ess_data$Parameter)[diagnostic_df$Parameter]
    }
    if ("ESS_bulk" %in% diagnostic && "ESS_bulk" %in% names(ess_data)) {
      # fmt: skip
      diagnostic_df$ESS_bulk <- stats::setNames(ess_data$ESS_bulk, ess_data$Parameter)[diagnostic_df$Parameter]
    }
  }

  if ("MCSE" %in% diagnostic) {
    diagnostic_df$MCSE <- mcse(posterior, effects = effects)$MCSE
  }

  if ("Rhat" %in% diagnostic) {
    s <- as.data.frame(rstan::summary(posterior)$summary)
    diagnostic_df$Rhat <- s[rownames(s) %in% all_params, ]$Rhat
  }

  # Remove columns with all Nans
  diagnostic_df <- diagnostic_df[!sapply(diagnostic_df, function(x) all(is.na(x)))]

  # Select rows
  diagnostic_df[diagnostic_df$Parameter %in% params, ]
}


#' @export
diagnostic_posterior.CmdStanFit <- function(
  posterior,
  diagnostic = "all",
  parameters = NULL,
  ...
) {
  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS", "Rhat", "MCSE")
  }

  insight::check_if_installed("posterior")
  insight::check_if_installed("cmdstanr")

  pars <- insight::find_parameters(posterior, flatten = TRUE)
  if (!is.null(parameters)) {
    pars <- pars[!grepl(parameters, pars)]
  }

  draws <- posterior$draws(format = "draws_df", variables = pars)

  out <- posterior::summarize_draws(
    draws,
    posterior::default_convergence_measures(),
    MCSE = posterior::mcse_mean
  )
  out <- datawizard::data_rename(
    as.data.frame(out),
    c(
      Parameter = "variable",
      ESS = "ess_bulk",
      ESS_tail = "ess_tail",
      Rhat = "rhat"
    )
  )

  out[, c("Parameter", diagnostic), drop = FALSE]
}


#' @export
diagnostic_posterior.blavaan <- function(posterior, diagnostic = "all", ...) {
  # Find parameters
  params <- suppressWarnings(insight::find_parameters(posterior, flatten = TRUE))

  out <- data.frame(Parameter = params)

  # If no diagnostic
  if (is.null(diagnostic)) {
    return(out)
  }

  diagnostic <- match.arg(
    diagnostic,
    c("ESS", "Rhat", "MCSE", "all"),
    several.ok = TRUE
  )
  if ("all" %in% diagnostic) {
    diagnostic <- c("ESS", "Rhat", "MCSE")
  } else {
    diagnostic <- diagnostic
    if ("Rhat" %in% diagnostic) diagnostic <- c(diagnostic, "khat")
  }

  # Get indices
  if ("Rhat" %in% diagnostic) {
    insight::check_if_installed("blavaan")

    Rhat <- blavaan::blavInspect(posterior, what = "psrf")
    Rhat <- data.frame(
      Parameter = colnames(insight::get_parameters(posterior)),
      Rhat = Rhat
    )
    out <- merge(out, Rhat, by = "Parameter", all = TRUE)
  }

  if ("ESS" %in% diagnostic) {
    ESS <- effective_sample(posterior)
    out <- merge(out, ESS, by = "Parameter", all = TRUE)
  }

  if ("MCSE" %in% diagnostic) {
    MCSE <- mcse(posterior)
    out <- merge(out, MCSE, by = "Parameter", all = TRUE)
  }

  unique(out)
}

Try the bayestestR package in your browser

Any scripts or data that you put into this service are public.

bayestestR documentation built on May 22, 2026, 1:06 a.m.