R/methods.R

Defines functions nobs.choicer_hb vcov.choicer_hb coef.choicer_hb print.summary.choicer_hb summary.choicer_hb print.choicer_hb blp.choicer_nl diversion_ratios.choicer_nl elasticities.choicer_nl predict.choicer_nl blp.choicer_mxl diversion_ratios.choicer_mxl elasticities.choicer_mxl blp.choicer_mnl diversion_ratios.choicer_mnl elasticities.choicer_mnl blp diversion_ratios elasticities apply_mxl_delta_method predict.choicer_mxl predict.choicer_mnl print_footer print_se_weighting print_coef_table print.summary.choicer_mnp print_bayes_coef_table summary.choicer_mnp build_bayes_coef_table nobs.choicer_mnp vcov.choicer_mnp coef.choicer_mnp print.choicer_mnp print.summary.choicer_nl summary.choicer_nl print.summary.choicer_mxl summary.choicer_mxl print.summary.choicer_mnl summary.choicer_mnl build_coef_table significance_code nobs.choicer_fit logLik.choicer_fit wesml_vcov.choicer_nl wesml_vcov.choicer_mnl wesml_vcov.choicer_mxl .wesml_vcov_impl wesml_vcov vcov.choicer_fit coef.choicer_fit print.choicer_fit model_display_name

Documented in blp blp.choicer_mnl blp.choicer_mxl blp.choicer_nl coef.choicer_fit coef.choicer_hb coef.choicer_mnp diversion_ratios diversion_ratios.choicer_mnl diversion_ratios.choicer_mxl diversion_ratios.choicer_nl elasticities elasticities.choicer_mnl elasticities.choicer_mxl elasticities.choicer_nl logLik.choicer_fit nobs.choicer_fit nobs.choicer_hb nobs.choicer_mnp predict.choicer_mnl predict.choicer_mxl predict.choicer_nl print.choicer_fit print.choicer_hb print.choicer_mnp print.summary.choicer_hb print.summary.choicer_mnl print.summary.choicer_mnp print.summary.choicer_mxl print.summary.choicer_nl summary.choicer_hb summary.choicer_mnl summary.choicer_mnp summary.choicer_mxl summary.choicer_nl vcov.choicer_fit vcov.choicer_hb vcov.choicer_mnp wesml_vcov wesml_vcov.choicer_mnl wesml_vcov.choicer_mxl wesml_vcov.choicer_nl

# S3 methods for choicer_fit objects

# Display names for model types
model_display_name <- function(model) {
  switch(model,
    mnl = "Multinomial Logit (MNL)",
    mxl = "Mixed Logit (MXL)",
    nl  = "Nested Logit (NL)",
    mnp = "Bayesian Multinomial Probit (MNP)",
    hmnl = "Hierarchical Bayesian Multinomial Logit (HMNL)",
    hmnp = "Hierarchical Bayesian Multinomial Probit (HMNP)",
    model
  )
}

# --- print -------------------------------------------------------------------

#' Print a choicer_fit object
#'
#' Prints a brief summary of the fitted model.
#'
#' @param x A choicer_fit object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' print(fit)
#' }
#' @export
print.choicer_fit <- function(x, ...) {
  cat(model_display_name(x$model), "model\n")
  cat("  N obs:", x$nobs, " | Parameters:", x$n_params, "\n")
  cat("  Log-likelihood:", format(x$loglik, digits = 6), "\n")
  cat("  AIC:", format(-2 * x$loglik + 2 * x$n_params, digits = 6), "\n")
  if (!is.na(x$convergence)) {
    cat("  Convergence:", x$convergence, "(", x$message, ")\n")
  }
  invisible(x)
}

# --- coef --------------------------------------------------------------------

#' Extract coefficients from a choicer_fit object
#'
#' @param object A choicer_fit object.
#' @param ... Additional arguments (ignored).
#' @returns Named numeric vector of estimated coefficients.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' coef(fit)
#' }
#' @export
coef.choicer_fit <- function(object, ...) {
  object$coefficients
}

# --- vcov --------------------------------------------------------------------

#' Extract variance-covariance matrix from a choicer_fit object
#'
#' With no arguments, returns the variance-covariance matrix implied by the
#' fit's own \code{se_method} (triggering lazy computation if needed). Passing
#' \code{type} recomputes a different variance estimator post hoc from the
#' stored data — no refit needed (requires \code{keep_data = TRUE}):
#' \describe{
#'   \item{\code{"hessian"}}{Inverse of the analytical negated Hessian.}
#'   \item{\code{"bhhh"}}{Inverse of the BHHH/OPG information
#'     \eqn{\sum_i w_i s_i s_i'}.}
#'   \item{\code{"robust"}}{Huber-White sandwich
#'     \eqn{A^{-1} (\sum_i w_i^2 s_i s_i') A^{-1}} — also the valid WESML
#'     variance under choice-based weighting.}
#'   \item{\code{"cluster"}}{Cluster-robust sandwich
#'     \eqn{A^{-1} (\sum_g g_g g_g') A^{-1}} with
#'     \eqn{g_g = \sum_{i \in g} w_i s_i} the within-cluster sum of weighted
#'     scores. Requires \code{cluster} (or a fit made with
#'     \code{cluster_col}). No small-sample correction is applied.}
#' }
#' Here \eqn{i} indexes \emph{choice situations}. For repeated choices by the
#' same decision maker (panel data), cluster on the decision maker.
#'
#' Note (mixed logit): clustering repairs the \emph{inference}, not the
#' \emph{estimand}. \code{run_mxlogit()} treats each choice situation as an
#' independent draw from the mixing distribution (a cross-sectional MSL
#' likelihood, not the panel product form), so on panel data the point
#' estimates target that cross-sectional model; \code{type = "cluster"} makes
#' their standard errors robust to within-person dependence but does not turn
#' the fit into a panel mixed logit. For panel random coefficients use
#' \code{\link{run_hmnlogit}} (\code{person_col}).
#'
#' @param object A choicer_fit object.
#' @param type \code{NULL} (default; return the as-fitted vcov) or one of
#'   \code{"hessian"}, \code{"bhhh"}, \code{"robust"}, \code{"cluster"}.
#' @param cluster Cluster labels for \code{type = "cluster"}, one per choice
#'   situation. Alignment to the prepared (id-sorted) choice situations is
#'   handled as follows:
#'   \itemize{
#'     \item \strong{Named} (recommended): names are matched against the
#'       choice-situation ids, so the vector is safe in any order. Build it by
#'       naming your per-situation labels with the id values.
#'     \item \strong{Unnamed}: taken to be in the prepared, id-sorted order; a
#'       warning flags that assumption. A vector of per-alternative (row-level)
#'       length is rejected.
#'   }
#'   Defaults to the labels stored at fit time via \code{cluster_col} (already
#'   aligned). Supplying \code{cluster} without \code{type} implies
#'   \code{type = "cluster"}. The safest route is to pass \code{cluster_col=}
#'   at fit time, which sidesteps post-hoc alignment entirely.
#' @param ... Additional arguments (ignored).
#' @returns Named variance-covariance matrix, or NULL if unavailable.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, person := rep(1:10, each = 5)[id]]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' vcov(fit)                          # as fitted (hessian)
#' vcov(fit, type = "robust")         # Huber-White, post hoc
#' # named by situation id -> safe regardless of order
#' cl <- dt[, person[1L], by = id]
#' vcov(fit, type = "cluster", cluster = setNames(cl$V1, cl$id))
#' }
#' @export
vcov.choicer_fit <- function(object, type = NULL, cluster = NULL, ...) {
  if (is.null(type)) {
    if (is.null(cluster)) {
      object <- ensure_vcov(object)
      return(object$vcov)
    }
    type <- "cluster"
  }
  type <- match.arg(type, c("hessian", "bhhh", "robust", "cluster"))
  res <- .assemble_score_vcov(object, type = type, cluster = cluster)
  if (!is.null(res$vcov)) {
    nms <- names(object$coefficients)
    rownames(res$vcov) <- nms
    colnames(res$vcov) <- nms
  }
  res$vcov
}

# --- wesml_vcov (robust / sandwich) -----------------------------------------

#' Robust (sandwich) variance for a weighted / choice-based logit fit
#'
#' Recomputes the robust Huber-White sandwich variance
#' \eqn{V = A^{-1} B A^{-1}} for a fitted multinomial (MNL), mixed (MXL) or
#' nested (NL) logit, where the bread
#' \eqn{A = \sum_i w_i (-H_i)} is the weighted negated Hessian and the meat
#' \eqn{B = \sum_i w_i^2 s_i s_i'} is the weight-squared outer product of the
#' per-individual scores. This is the appropriate variance under choice-based
#' (endogenous stratified) / WESML weighting, where the inverse-Hessian and the
#' ordinary BHHH variance are invalid. It can be called on any fitted model
#' (e.g. one estimated with \code{se_method = "hessian"}) to obtain robust
#' standard errors post hoc, without refitting.
#'
#' If the stored weights are uniform (all equal), a warning is emitted: the
#' returned variance is then the ordinary robust (Huber-White) variance, not a
#' WESML-weighted variance. Refit with WESML weights for a choice-based-sampling
#' correction.
#'
#' @param object A fitted \code{choicer_mnl}, \code{choicer_mxl} or
#'   \code{choicer_nl} object (requires \code{keep_data = TRUE}).
#' @param type Either \code{"vcov"} (default) to return the variance-covariance
#'   matrix or \code{"se"} to return the standard-error vector.
#' @param ... Unused.
#' @returns A variance-covariance matrix (\code{type = "vcov"}) or a named
#'   numeric vector of standard errors (\code{type = "se"}), in the raw
#'   parameter space.
#' @seealso \code{\link{wesml_weights}}, \code{\link{sample_by_choice}},
#'   \code{\link{run_mxlogit}}
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(1)
#' N <- 200L; J <- 3L
#' dt <- data.table(id = rep(seq_len(N), each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := as.integer(seq_len(.N) == sample.int(.N, 1L)), by = id]
#' fit <- run_mxlogit(dt, "id", "alt", "choice", "x1", "w1", S = 50L)
#' wesml_vcov(fit, "se")
#' }
#' @export
wesml_vcov <- function(object, ...) UseMethod("wesml_vcov")

# Shared implementation for the wesml_vcov methods. The model-specific dispatch
# (MNL/NL/MXL) is handled inside compute_sandwich_vcov(); the surrounding
# stored-data check, uniform-weight warning, and naming logic are identical.
.wesml_vcov_impl <- function(object, type = c("vcov", "se")) {
  type <- match.arg(type)
  if (is.null(object[["data"]])) {
    stop("wesml_vcov() needs the stored data; refit with keep_data = TRUE.")
  }
  w <- object[["data"]]$weights
  if (!is.null(w) && length(unique(w)) == 1L) {
    warning("Stored weights are uniform; wesml_vcov() returns the ordinary robust ",
            "(Huber-White) variance, not a WESML-weighted variance. Refit with WESML ",
            "weights for a choice-based-sampling correction.", call. = FALSE)
  }
  res <- compute_sandwich_vcov(object)
  nms <- names(object$coefficients)
  if (!is.null(res$vcov)) {
    rownames(res$vcov) <- nms
    colnames(res$vcov) <- nms
  }
  if (!is.null(res$se)) names(res$se) <- nms
  if (type == "se") res$se else res$vcov
}

#' @rdname wesml_vcov
#' @export
wesml_vcov.choicer_mxl <- function(object, type = c("vcov", "se"), ...) {
  .wesml_vcov_impl(object, type)
}

#' @rdname wesml_vcov
#' @export
wesml_vcov.choicer_mnl <- function(object, type = c("vcov", "se"), ...) {
  .wesml_vcov_impl(object, type)
}

#' @rdname wesml_vcov
#' @export
wesml_vcov.choicer_nl <- function(object, type = c("vcov", "se"), ...) {
  .wesml_vcov_impl(object, type)
}

# --- logLik ------------------------------------------------------------------

#' Extract log-likelihood from a choicer_fit object
#'
#' Returns a logLik object, which enables AIC() and BIC() automatically.
#'
#' @param object A choicer_fit object.
#' @param ... Additional arguments (ignored).
#' @returns A logLik object with df and nobs attributes.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' logLik(fit)
#' AIC(fit)
#' BIC(fit)
#' }
#' @export
logLik.choicer_fit <- function(object, ...) {
  val <- object$loglik
  attr(val, "df") <- object$n_params
  attr(val, "nobs") <- object$nobs
  class(val) <- "logLik"
  val
}

# --- nobs --------------------------------------------------------------------

#' Extract number of observations from a choicer_fit object
#'
#' @param object A choicer_fit object.
#' @param ... Additional arguments (ignored).
#' @returns Integer number of choice situations.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' nobs(fit)
#' }
#' @export
nobs.choicer_fit <- function(object, ...) {
  object$nobs
}

# --- summary: shared helpers -------------------------------------------------

#' Significance codes for p-values
#' @noRd
significance_code <- function(p) {
  if (is.na(p)) return("")
  if (p < 0.001) return("***")
  if (p < 0.01) return("**")
  if (p < 0.05) return("*")
  ""
}

#' Build the standard coefficient table
#' @noRd
build_coef_table <- function(estimates, se, param_names) {
  zval <- estimates / se
  pval <- 2 * (1 - pnorm(abs(zval)))

  data.frame(
    Estimate  = estimates,
    Std_Error = se,
    z_value   = zval,
    Pr_z      = pval,
    Signif    = vapply(pval, significance_code, character(1)),
    row.names = param_names,
    stringsAsFactors = FALSE
  )
}

# --- summary: MNL ------------------------------------------------------------

#' Summary for multinomial logit model
#'
#' Computes and returns a coefficient summary table with standard errors,
#' z-values, p-values, and significance codes. Triggers lazy Hessian
#' computation if standard errors have not been computed yet.
#'
#' @param object A choicer_mnl object.
#' @param gof Logical; compute goodness-of-fit measures (McFadden R-squared,
#'   hit rate) for the summary footer. Involves an in-sample prediction pass
#'   (for mixed logit, a full simulation over draws); set to FALSE to skip.
#' @param ... Additional arguments (ignored).
#' @returns A summary.choicer_mnl object (list with coefficients table and
#'   metadata, including a `gof` element with goodness-of-fit measures from
#'   \code{\link{gof}}; its fields are NA when the model was fitted with
#'   \code{keep_data = FALSE}).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' summary(fit)
#' }
#' @export
summary.choicer_mnl <- function(object, gof = TRUE, ...) {
  object <- ensure_vcov(object)

  coef_table <- build_coef_table(
    estimates = object$coefficients,
    se = object$se %||% rep(NA_real_, length(object$coefficients)),
    param_names = names(object$coefficients)
  )

  structure(
    list(
      model = object$model,
      coefficients = coef_table,
      loglik = object$loglik,
      nobs = object$nobs,
      n_params = object$n_params,
      convergence = object$convergence,
      message = object$message,
      elapsed_time = object$optimizer$elapsed_time,
      se_method = object$se_method %||% "hessian",
      weighting = object$choice_sampling$scheme,
      weights_applied = object$choice_sampling$weights_applied,
      gof = if (isTRUE(gof)) {
        # calling gof() here is safe: R skips the logical binding when
        # resolving a function call
        tryCatch(suppressMessages(gof(object)), error = function(e) NULL)
      }
    ),
    class = "summary.choicer_mnl"
  )
}

#' Print summary for multinomial logit model
#' @param x A summary.choicer_mnl object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' print(summary(fit))
#' }
#' @export
print.summary.choicer_mnl <- function(x, ...) {
  cat(model_display_name(x$model), "model\n\n")
  print_coef_table(x$coefficients)
  cat("\n")
  print_se_weighting(x)
  print_footer(x)
  invisible(x)
}

# --- summary: MXL ------------------------------------------------------------

#' Summary for mixed logit model
#'
#' Computes coefficient summary with delta-method transformation for variance
#' parameters (Cholesky to covariance scale) and log-normal mean parameters.
#' Triggers lazy Hessian computation if standard errors have not been computed yet.
#'
#' @param object A choicer_mxl object.
#' @param gof Logical; compute goodness-of-fit measures (McFadden R-squared,
#'   hit rate) for the summary footer. Involves an in-sample prediction pass
#'   (for mixed logit, a full simulation over draws); set to FALSE to skip.
#' @param ... Additional arguments (ignored).
#' @returns A summary.choicer_mxl object (includes a `gof` element with
#'   goodness-of-fit measures from \code{\link{gof}}; its fields are NA when
#'   the model was fitted with \code{keep_data = FALSE}).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' summary(fit)
#' }
#' @export
summary.choicer_mxl <- function(object, gof = TRUE, ...) {
  object <- ensure_vcov(object)

  est <- object$coefficients
  se <- object$se %||% rep(NA_real_, length(est))

  # Apply delta method for log-normal mu and Cholesky -> Sigma
  if (!is.null(object$vcov) && !is.null(object$param_map)) {
    delta_result <- apply_mxl_delta_method(
      est_theta = est,
      se = se,
      vcov_mat = object$vcov,
      param_map = object$param_map,
      rc_dist = object$rc_dist,
      rc_correlation = object$rc_correlation,
      rc_mean = object$rc_mean
    )
    est <- delta_result$estimates
    se <- delta_result$se
  }

  # Build display names: L_ij -> Sigma_ij, Mu_x -> exp(Mu_x) for log-normal
  display_names <- names(object$coefficients)
  if (!is.null(object$param_map$sigma)) {
    idx_sigma <- object$param_map$sigma
    K_w <- length(object$rc_dist)
    if (object$rc_correlation) {
      sigma_display <- character(length(idx_sigma))
      k <- 1
      for (i in seq_len(K_w)) {
        for (j in seq_len(i)) {
          sigma_display[k] <- sprintf("Sigma_%d%d", i, j)
          k <- k + 1
        }
      }
    } else {
      sigma_display <- paste0("Sigma_", seq_len(K_w), seq_len(K_w))
    }
    display_names[idx_sigma] <- sigma_display
  }
  if (object$rc_mean && !is.null(object$param_map$mu)) {
    K_w <- length(object$rc_dist)
    for (k in seq_len(K_w)) {
      if (object$rc_dist[k] == 1) {
        idx <- object$param_map$mu[k]
        display_names[idx] <- paste0("exp(", display_names[idx], ")")
      }
    }
  }

  coef_table <- build_coef_table(
    estimates = est,
    se = se,
    param_names = display_names
  )

  structure(
    list(
      model = object$model,
      coefficients = coef_table,
      loglik = object$loglik,
      nobs = object$nobs,
      n_params = object$n_params,
      convergence = object$convergence,
      message = object$message,
      elapsed_time = object$optimizer$elapsed_time,
      sigma = object$sigma,
      se_method = object$se_method %||% "hessian",
      weighting = object$choice_sampling$scheme,
      weights_applied = object$choice_sampling$weights_applied,
      gof = if (isTRUE(gof)) {
        # calling gof() here is safe: R skips the logical binding when
        # resolving a function call
        tryCatch(suppressMessages(gof(object)), error = function(e) NULL)
      }
    ),
    class = "summary.choicer_mxl"
  )
}

#' Print summary for mixed logit model
#' @param x A summary.choicer_mxl object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' print(summary(fit))
#' }
#' @export
print.summary.choicer_mxl <- function(x, ...) {
  cat(model_display_name(x$model), "model\n\n")
  print_coef_table(x$coefficients)
  cat("\n")
  if (!is.null(x$sigma)) {
    cat("Random coefficient covariance (Sigma):\n")
    print(x$sigma)
    cat("\n")
  }
  print_se_weighting(x)
  print_footer(x)
  invisible(x)
}

# --- summary: NL -------------------------------------------------------------

#' Summary for nested logit model
#'
#' Triggers lazy Hessian computation if standard errors have not been computed yet.
#'
#' @param object A choicer_nl object.
#' @param gof Logical; compute goodness-of-fit measures (McFadden R-squared,
#'   hit rate) for the summary footer. Involves an in-sample prediction pass
#'   (for mixed logit, a full simulation over draws); set to FALSE to skip.
#' @param ... Additional arguments (ignored).
#' @returns A summary.choicer_nl object (includes a `gof` element with
#'   goodness-of-fit measures from \code{\link{gof}}; its fields are NA when
#'   the model was fitted with \code{keep_data = FALSE}).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, nest := ifelse(alt <= 2, "A", "B")]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = c("x1", "x2"), nest_col = "nest"
#' )
#' summary(fit)
#' }
#' @export
summary.choicer_nl <- function(object, gof = TRUE, ...) {
  object <- ensure_vcov(object)

  coef_table <- build_coef_table(
    estimates = object$coefficients,
    se = object$se %||% rep(NA_real_, length(object$coefficients)),
    param_names = names(object$coefficients)
  )

  structure(
    list(
      model = object$model,
      coefficients = coef_table,
      loglik = object$loglik,
      nobs = object$nobs,
      n_params = object$n_params,
      convergence = object$convergence,
      message = object$message,
      elapsed_time = object$optimizer$elapsed_time,
      se_method = object$se_method %||% "hessian",
      weighting = object$choice_sampling$scheme,
      weights_applied = object$choice_sampling$weights_applied,
      gof = if (isTRUE(gof)) {
        # calling gof() here is safe: R skips the logical binding when
        # resolving a function call
        tryCatch(suppressMessages(gof(object)), error = function(e) NULL)
      }
    ),
    class = "summary.choicer_nl"
  )
}

#' Print summary for nested logit model
#' @param x A summary.choicer_nl object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, nest := ifelse(alt <= 2, "A", "B")]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = c("x1", "x2"), nest_col = "nest"
#' )
#' print(summary(fit))
#' }
#' @export
print.summary.choicer_nl <- function(x, ...) {
  cat(model_display_name(x$model), "model\n\n")
  print_coef_table(x$coefficients)
  cat("\n")
  print_se_weighting(x)
  print_footer(x)
  invisible(x)
}

# --- MNP (Bayesian) methods ----------------------------------------------------
# choicer_mnp is a posterior-draws object and intentionally does not inherit
# from choicer_fit: there is no log-likelihood, convergence code, or lazy
# Hessian. All summaries are computed from the identified draws
# (beta / sqrt(sigma_11), Sigma / sigma_11).

#' Print a choicer_mnp object
#'
#' Prints a brief summary of the fitted Bayesian multinomial probit model.
#'
#' @param x A choicer_mnp object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' print(fit)
#' }
#' @export
print.choicer_mnp <- function(x, ...) {
  cat(model_display_name(x$model), "model\n")
  cat("  N obs:", x$nobs, " | Parameters:", x$n_params, "\n")
  cat("  Posterior draws kept:", x$mcmc$R_keep,
      sprintf("(R = %d, burn = %d, thin = %d)\n", x$mcmc$R, x$mcmc$burn, x$mcmc$thin))
  cat("  Base alternative:", as.character(x$base_alt), "\n")
  cat("  Estimates are posterior means of identified parameters",
      "(beta / sqrt(sigma_11)).\n")
  invisible(x)
}

#' Extract coefficients from a choicer_mnp object
#'
#' Returns the posterior means of the identified coefficients
#' (\eqn{\beta / \sqrt{\sigma_{11}}}, computed per draw).
#'
#' @param object A choicer_mnp object.
#' @param ... Additional arguments (ignored).
#' @returns Named numeric vector of posterior means.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' coef(fit)
#' }
#' @export
coef.choicer_mnp <- function(object, ...) {
  object$coefficients
}

#' Extract variance-covariance matrix from a choicer_mnp object
#'
#' Returns the posterior covariance matrix of the identified coefficient
#' draws (computed eagerly at fit time; no Hessian is involved).
#'
#' @param object A choicer_mnp object.
#' @param ... Additional arguments (ignored).
#' @returns Named posterior covariance matrix.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' vcov(fit)
#' }
#' @export
vcov.choicer_mnp <- function(object, ...) {
  object$vcov
}

#' Extract number of observations from a choicer_mnp object
#'
#' @param object A choicer_mnp object.
#' @param ... Additional arguments (ignored).
#' @returns Integer number of choice situations.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' nobs(fit)
#' }
#' @export
nobs.choicer_mnp <- function(object, ...) {
  object$nobs
}

#' Build a Bayesian posterior summary table
#'
#' Posterior mean, SD, and equal-tailed credible interval bounds, one row per
#' column of the draw matrix. No z- or p-values: those are frequentist
#' quantities with no role in a posterior summary.
#' @noRd
build_bayes_coef_table <- function(draws, prob = 0.95) {
  alpha <- (1 - prob) / 2
  qs <- t(apply(draws, 2, stats::quantile, probs = c(alpha, 0.5, 1 - alpha)))
  data.frame(
    Mean   = colMeans(draws),
    SD     = apply(draws, 2, stats::sd),
    CI_lo  = qs[, 1],
    Median = qs[, 2],
    CI_hi  = qs[, 3],
    row.names = colnames(draws),
    stringsAsFactors = FALSE
  )
}

#' Summary for Bayesian multinomial probit model
#'
#' Posterior summaries (mean, SD, equal-tailed credible interval) of the
#' identified coefficient and covariance draws.
#'
#' @param object A choicer_mnp object.
#' @param prob Probability mass of the equal-tailed credible interval
#'   (default 0.95).
#' @param ... Additional arguments (ignored).
#' @returns A summary.choicer_mnp object (list with coefficient and Sigma
#'   posterior tables plus metadata).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' summary(fit)
#' }
#' @export
summary.choicer_mnp <- function(object, prob = 0.95, ...) {
  if (!is.numeric(prob) || length(prob) != 1L || prob <= 0 || prob >= 1) {
    stop("prob must be a single number in (0, 1).")
  }

  structure(
    list(
      model = object$model,
      coefficients = build_bayes_coef_table(object$draws$beta, prob),
      sigma_table = build_bayes_coef_table(object$draws$sigma, prob),
      sigma = object$sigma,
      prob = prob,
      base_alt = object$base_alt,
      nobs = object$nobs,
      n_params = object$n_params,
      mcmc = object$mcmc,
      elapsed_time = object$sampler$elapsed_time
    ),
    class = "summary.choicer_mnp"
  )
}

#' Print a formatted Bayesian posterior summary table
#' @noRd
print_bayes_coef_table <- function(coef_table, prob) {
  param_width <- max(nchar(rownames(coef_table)), nchar("Parameter"))
  alpha <- (1 - prob) / 2
  lo_lab <- sprintf("%g%%", 100 * alpha)
  hi_lab <- sprintf("%g%%", 100 * (1 - alpha))

  cat(sprintf(
    "%-*s  %10s %10s %10s %10s %10s\n",
    param_width, "Parameter", "Mean", "SD", lo_lab, "Median", hi_lab
  ))

  for (i in seq_len(nrow(coef_table))) {
    cat(sprintf(
      "%-*s  %10.6f %10.6f %10.6f %10.6f %10.6f\n",
      param_width,
      rownames(coef_table)[i],
      coef_table$Mean[i],
      coef_table$SD[i],
      coef_table$CI_lo[i],
      coef_table$Median[i],
      coef_table$CI_hi[i]
    ))
  }
}

#' Print summary for Bayesian multinomial probit model
#' @param x A summary.choicer_mnp object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 100; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnprobit(dt, "id", "alt", "choice", c("x1", "x2"),
#'                     mcmc = list(R = 300, burn = 100))
#' print(summary(fit))
#' }
#' @export
print.summary.choicer_mnp <- function(x, ...) {
  cat(model_display_name(x$model), "model\n\n")
  print_bayes_coef_table(x$coefficients, x$prob)
  cat("\n")
  cat("Covariance of utility differences (Sigma, identified scale):\n")
  print_bayes_coef_table(x$sigma_table, x$prob)
  cat("\nPosterior mean Sigma:\n")
  print(round(x$sigma, 6))
  cat("\n")
  cat("Base alternative:", as.character(x$base_alt), "\n")
  cat("Draws kept:", x$mcmc$R_keep,
      sprintf("(R = %d, burn = %d, thin = %d, seed = %d)\n",
              x$mcmc$R, x$mcmc$burn, x$mcmc$thin, as.integer(x$mcmc$seed)))
  cat("N:", x$nobs, " | Parameters:", x$n_params, "\n")
  if (!is.null(x$elapsed_time)) {
    cat("Sampling time:", round(x$elapsed_time, 2), "s\n")
  }
  cat("Identification: per-draw normalization by sigma_11",
      "(McCulloch-Rossi 1994).\n")
  invisible(x)
}

# --- Shared printing helpers -------------------------------------------------

#' Print a formatted coefficient table
#' @noRd
print_coef_table <- function(coef_table) {
  # Column widths
  param_width <- max(nchar(rownames(coef_table)), nchar("Parameter"))

  cat(sprintf(
    "%-*s  %10s %10s %8s %9s  %s\n",
    param_width, "Parameter", "Estimate", "Std.Error", "z-value", "Pr(>|z|)", ""
  ))

  for (i in seq_len(nrow(coef_table))) {
    cat(sprintf(
      "%-*s  %10.6f %10.6f %8.4f %9.2e  %s\n",
      param_width,
      rownames(coef_table)[i],
      coef_table$Estimate[i],
      coef_table$Std_Error[i],
      coef_table$z_value[i],
      coef_table$Pr_z[i],
      coef_table$Signif[i]
    ))
  }

  cat("---\nSignif. codes:  '***' 0.001 '**' 0.01 '*' 0.05\n")
}

#' Print model footer (log-likelihood, AIC, timing)
#' @noRd
# Print the Std. Errors method label and (optional) WESML weighting line.
# Shared by the MNL / MXL / NL summary print methods.
print_se_weighting <- function(x) {
  cat("Std. Errors:", switch(
    x$se_method %||% "hessian",
    bhhh = "BHHH (OPG)",
    sandwich = "Sandwich (robust)",
    cluster = "Cluster-robust sandwich",
    numeric = "Numerical Hessian (finite differences)",
    "Analytical Hessian"
  ), "\n")
  if (!is.null(x$weighting)) {
    # Backward-compat: when weights_applied is absent (older fits) treat as applied.
    applied <- !isFALSE(x$weights_applied)
    cat("Weighting:",
        if (identical(x$weighting, "wesml")) {
          if (applied) {
            "WESML choice-based"
          } else {
            "WESML provenance present but NOT applied (fit is unweighted)"
          }
        } else {
          "user-supplied"
        },
        "\n")
  }
}

print_footer <- function(x) {
  cat("Log-likelihood:", format(x$loglik, digits = 6), "\n")
  aic <- -2 * x$loglik + 2 * x$n_params
  bic <- -2 * x$loglik + log(x$nobs) * x$n_params
  cat("AIC:", format(aic, digits = 6), " | BIC:", format(bic, digits = 6), "\n")
  print_gof_lines(x$gof)
  cat("N:", x$nobs, " | Parameters:", x$n_params, "\n")
  if (!is.null(x$elapsed_time)) {
    cat("Optimization time:", round(x$elapsed_time, 2), "s\n")
  }
  if (!is.na(x$convergence)) {
    cat("Convergence:", x$convergence, "(", x$message, ")\n")
  }
}

# --- predict: MNL ------------------------------------------------------------

#' Predict from a multinomial logit model
#'
#' Computes choice probabilities or aggregate market shares, either for the
#' data used at fit time (default) or for counterfactual `newdata`.
#'
#' @param object A choicer_mnl object.
#' @param type One of "probabilities" (individual-level choice probabilities)
#'   or "shares" (aggregate market shares).
#' @param newdata Optional data for counterfactual prediction. Either:
#'   * a data.frame in the same long format used at fit time (one row per
#'     id-alternative pair, with the fit-time id, alternative, and covariate
#'     columns; a choice column is not required). Alternative labels must have
#'     been seen at fit time; per-id subsets of alternatives are allowed.
#'   * a list with elements `X`, `alt_idx`, `M` (and optionally `weights`)
#'     matching the layout of `object$data` — the "modified design matrix"
#'     path for policy simulation (e.g., perturb a column of `object$data$X`).
#'     `alt_idx` must use the fit-time integer codes from `object$alt_mapping`.
#'
#'   When `NULL` (default), the data stored at fit time is used (requires
#'   `keep_data = TRUE`).
#' @param weights Optional numeric vector with one weight per choice situation,
#'   used for `type = "shares"` aggregation. For a data.frame `newdata`,
#'   supply one weight per id in order of first appearance in `newdata`
#'   (weights are realigned internally to the sorted row order). Defaults to
#'   equal weights. Ignored when `newdata` is `NULL` (the stored fit weights
#'   apply).
#' @param ... Additional arguments (ignored).
#' @returns For "probabilities": a list with `choice_prob` and `utility` vectors.
#'   For "shares": a named numeric vector of market shares per alternative.
#'   With a data.frame `newdata`, rows are ordered by id, then by fit-time
#'   alternative code (`alt_int` in `object$alt_mapping`).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' predict(fit, type = "shares")
#' predict(fit, type = "probabilities")
#'
#' # Counterfactual: increase x1 for alternative 2
#' dt_cf <- copy(dt)[alt == 2, x1 := x1 + 1]
#' predict(fit, type = "shares", newdata = dt_cf)
#' }
#' @export
predict.choicer_mnl <- function(object, type = c("probabilities", "shares"),
                                newdata = NULL, weights = NULL, ...) {
  type <- match.arg(type)

  if (is.null(newdata)) {
    if (is.null(object[["data"]])) {
      stop("Prediction requires stored data. Refit with keep_data = TRUE.")
    }
    d <- object[["data"]]
  } else {
    d <- resolve_predict_newdata(object, newdata, weights = weights)
  }

  theta <- object$coefficients

  if (type == "probabilities") {
    mnl_predict(
      theta = theta,
      X = d$X,
      alt_idx = d$alt_idx,
      M = d$M,
      use_asc = object$use_asc,
      include_outside_option = object$include_outside_option
    )
  } else {
    mnl_predict_shares(
      theta = theta,
      X = d$X,
      alt_idx = d$alt_idx,
      M = d$M,
      weights = d$weights,
      use_asc = object$use_asc,
      include_outside_option = object$include_outside_option
    )
  }
}

# --- predict: MXL ------------------------------------------------------------

#' Predict from a mixed logit model
#'
#' Computes simulated choice probabilities or aggregate market shares using
#' deterministic Halton draws, either for the data used at fit time (default)
#' or for counterfactual `newdata`.
#'
#' @param object A choicer_mxl object.
#' @param type Either "probabilities" (per-observation simulated choice
#'   probabilities) or "shares" (aggregate simulated market shares).
#' @param newdata Optional data for counterfactual prediction. Either:
#'   * a data.frame in the same long format used at fit time (one row per
#'     id-alternative pair, with the fit-time id, alternative, fixed-coefficient,
#'     and random-coefficient columns; a choice column is not required).
#'     Alternative labels must have been seen at fit time; per-id subsets of
#'     alternatives are allowed.
#'   * a list with elements `X`, `W`, `alt_idx`, `M` (and optionally
#'     `weights`) matching the layout of `object$data` — the "modified design
#'     matrix" path for policy simulation. `alt_idx` must use the fit-time
#'     integer codes from `object$alt_mapping`.
#'
#'   When `NULL` (default), the data stored at fit time is used (requires
#'   `keep_data = TRUE`). Halton draws are regenerated deterministically from
#'   `object$draws_info` with one block of draws per choice situation in
#'   `newdata`.
#' @param weights Optional numeric vector with one weight per choice situation,
#'   used for `type = "shares"` aggregation. For a data.frame `newdata`,
#'   supply one weight per id in order of first appearance in `newdata`
#'   (weights are realigned internally to the sorted row order). Defaults to
#'   equal weights. Ignored when `newdata` is `NULL` (the stored fit weights
#'   apply).
#' @param ... Additional arguments (ignored).
#' @returns For "probabilities": a list with `choice_prob` and `utility`
#'   vectors averaged across simulation draws. For "shares": a named numeric
#'   vector of simulated market shares per alternative. With a data.frame
#'   `newdata`, rows are ordered by id, then by fit-time alternative code
#'   (`alt_int` in `object$alt_mapping`).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' predict(fit, type = "shares")
#' predict(fit, type = "probabilities")
#' }
#' @export
predict.choicer_mxl <- function(object, type = c("probabilities", "shares"),
                                newdata = NULL, weights = NULL, ...) {
  type <- match.arg(type)

  if (is.null(newdata)) {
    if (is.null(object[["data"]]) || is.null(object$draws_info)) {
      stop("Prediction requires stored data and draws. ",
           "Refit with keep_data = TRUE.")
    }
    d <- object[["data"]]
    N_draws <- object$draws_info$N
  } else {
    if (is.null(object$draws_info)) {
      stop("Prediction with newdata requires stored draw metadata ",
           "('draws_info'). Refit to enable newdata prediction.")
    }
    d <- resolve_predict_newdata(object, newdata, weights = weights)
    N_draws <- d$N
  }

  # Resolve draws: in generate mode use empty placeholder + gen params.
  # Note: newdata with a different N is handled automatically by the C++ generator
  # (which uses (i-1)*S+s+1 regardless of N); for store mode we regenerate for N_draws.
  mode_pred <- (object$draws_info$mode) %||% "store"
  if (mode_pred == "store") {
    eta_draws        <- get_halton_normals(
      S   = object$draws_info$S,
      N   = N_draws,
      K_w = object$draws_info$K_w
    )
    gen_seed_arg     <- -1L
    gen_scramble_arg <- 1L
    gen_S_arg        <- 0L
  } else {
    eta_draws        <- array(0, dim = c(object$draws_info$K_w, 0L, 0L))
    gen_seed_arg     <- as.integer(object$draws_info$seed)
    gen_scramble_arg <- if (object$draws_info$scramble %in%
                             c("permuted", "owen")) 1L else 0L
    gen_S_arg        <- as.integer(object$draws_info$S)
  }

  args <- list(
    theta                  = object$coefficients,
    X                      = d$X,
    W                      = d$W,
    alt_idx                = d$alt_idx,
    M                      = d$M,
    eta_draws              = eta_draws,
    rc_dist                = object$rc_dist,
    rc_correlation         = object$rc_correlation,
    rc_mean                = object$rc_mean,
    use_asc                = object$use_asc,
    include_outside_option = object$include_outside_option,
    gen_seed               = gen_seed_arg,
    gen_scramble           = gen_scramble_arg,
    gen_S                  = gen_S_arg
  )

  if (type == "probabilities") {
    do.call(mxl_predict, args)
  } else {
    do.call(mxl_predict_shares, c(args, list(weights = d$weights)))
  }
}

# --- Delta method for MXL summary -------------------------------------------

#' Apply delta method for MXL variance parameters
#'
#' Transforms Cholesky parameters to covariance scale and
#' applies delta method for log-normal mu parameters.
#' Uses param_map for stable parameter indexing.
#'
#' @param est_theta Raw parameter estimates
#' @param se Raw standard errors
#' @param vcov_mat Variance-covariance matrix
#' @param param_map Named list with index vectors (beta, mu, sigma, asc)
#' @param rc_dist Integer vector of distribution types
#' @param rc_correlation Logical whether correlated
#' @param rc_mean Logical whether mu estimated
#' @returns List with transformed estimates and se
#' @noRd
apply_mxl_delta_method <- function(est_theta, se, vcov_mat,
                                   param_map, rc_dist, rc_correlation,
                                   rc_mean) {
  est <- est_theta
  se_out <- se

  K_w <- length(rc_dist)

  # Delta method for log-normal mu: exp(mu)
  if (rc_mean && !is.null(param_map$mu)) {
    idx_mu <- param_map$mu
    for (k in seq_len(K_w)) {
      if (rc_dist[k] == 1) {
        curr_idx <- idx_mu[k]
        mu_hat <- est[curr_idx]
        est[curr_idx] <- exp(mu_hat)
        se_out[curr_idx] <- exp(mu_hat) * se[curr_idx]
      }
    }
  }

  # Delta method for Cholesky -> Sigma
  if (!is.null(param_map$sigma)) {
    idx_sigma <- param_map$sigma
    L_params_hat <- est_theta[idx_sigma]
    Sigma_hat <- build_var_mat(L_params_hat, K_w, rc_correlation)

    if (rc_correlation) {
      est[idx_sigma] <- vech_row(Sigma_hat)
    } else {
      est[idx_sigma] <- diag(Sigma_hat)
    }

    J_mat <- jacobian_vech_Sigma(L_params_hat, K_w, rc_correlation)
    vcov_L <- vcov_mat[idx_sigma, idx_sigma]
    V_sigma <- J_mat %*% vcov_L %*% t(J_mat)
    se_out[idx_sigma] <- sqrt(diag(V_sigma))
  }

  list(estimates = est, se = se_out)
}

# --- Post-estimation generics ------------------------------------------------

#' Compute aggregate elasticities
#'
#' Computes a J x J matrix of aggregate elasticities. Entry (i, j) is the
#' percentage change in the probability of choosing alternative i when the
#' attribute of alternative j changes by 1\%.
#'
#' @param object A fitted model object.
#' @param ... Additional arguments passed to methods.
#' @returns A J x J elasticity matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' elasticities(fit, "x1")
#' }
#' @export
elasticities <- function(object, ...) UseMethod("elasticities")

#' Compute aggregate diversion ratios
#'
#' Computes a J x J matrix of diversion ratios. Entry (i, j) is the fraction
#' of demand lost by alternative j that is captured by alternative i when
#' alternative j becomes less attractive.
#'
#' @param object A fitted model object.
#' @param ... Additional arguments passed to methods.
#' @returns A J x J diversion ratio matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' diversion_ratios(fit)
#' }
#' @export
diversion_ratios <- function(object, ...) UseMethod("diversion_ratios")

#' BLP contraction mapping
#'
#' Finds the ASC (delta) parameters such that predicted market shares match
#' target shares, using the contraction mapping of Berry, Levinsohn, and
#' Pakes (1995) \doi{10.2307/2171802}.
#'
#' @param object A fitted model object.
#' @param target_shares Numeric vector of target market shares (length J).
#' @param ... Additional arguments passed to methods.
#' @returns Converged delta (ASC) vector.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' blp(fit, target_shares = rep(1/J, J))
#' }
#' @export
blp <- function(object, target_shares, ...) UseMethod("blp")

# --- elasticities: MNL -------------------------------------------------------

#' Elasticities for multinomial logit model
#'
#' @param object A \code{choicer_mnl} object fitted with \code{keep_data = TRUE}.
#' @param elast_var Variable for elasticity computation: a column name (character)
#'   or 1-based index into the design matrix X.
#' @param ... Additional arguments (ignored).
#' @returns A J x J elasticity matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' elasticities(fit, "x1")
#' }
#' @export
elasticities.choicer_mnl <- function(object, elast_var, ...) {
  if (is.null(object[["data"]])) {
    stop("elasticities() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]
  idx <- resolve_var_index(elast_var, colnames(d$X))

  mat <- mnl_elasticities_parallel(
    theta = object$coefficients,
    X = d$X,
    alt_idx = d$alt_idx,
    choice_idx = d$choice_idx,
    M = d$M,
    weights = d$weights,
    elast_var_idx = idx,
    use_asc = object$use_asc,
    include_outside_option = object$include_outside_option
  )

  label_matrix(mat, object$alt_mapping)
}

# --- diversion_ratios: MNL ---------------------------------------------------

#' Diversion ratios for multinomial logit model
#'
#' @param object A \code{choicer_mnl} object fitted with \code{keep_data = TRUE}.
#' @param ... Additional arguments (ignored).
#' @returns A J x J diversion ratio matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' diversion_ratios(fit)
#' }
#' @export
diversion_ratios.choicer_mnl <- function(object, ...) {
  if (is.null(object[["data"]])) {
    stop("diversion_ratios() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]

  mat <- mnl_diversion_ratios_parallel(
    theta = object$coefficients,
    X = d$X,
    alt_idx = d$alt_idx,
    M = d$M,
    weights = d$weights,
    use_asc = object$use_asc,
    include_outside_option = object$include_outside_option
  )

  label_matrix(mat, object$alt_mapping)
}

# --- blp: MNL ----------------------------------------------------------------

#' BLP contraction mapping for multinomial logit model
#'
#' @param object A \code{choicer_mnl} object fitted with \code{keep_data = TRUE}.
#' @param target_shares Numeric vector of target market shares.
#'   Length \code{J_inside} when no outside option, or \code{J_inside + 1}
#'   (with the outside option's share at index 1) when
#'   \code{include_outside_option = TRUE}.
#' @param delta_init Initial guess for delta (ASC) values. If \code{NULL},
#'   uses the estimated ASCs from the fitted model.
#' @param tol Convergence tolerance (default 1e-8).
#' @param max_iter Maximum iterations (default 1000).
#' @param ... Additional arguments (ignored).
#' @returns Converged delta (ASC) vector.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mnlogit(dt, "id", "alt", "choice", c("x1", "x2"))
#' blp(fit, target_shares = rep(1/J, J))
#' }
#' @export
blp.choicer_mnl <- function(object, target_shares, delta_init = NULL,
                            tol = 1e-8, max_iter = 1000, ...) {
  if (is.null(object[["data"]])) {
    stop("blp() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]
  pm <- object$param_map
  beta <- object$coefficients[pm$beta]

  J <- nrow(object$alt_mapping)
  if (is.null(delta_init)) {
    if (!is.null(pm$asc)) {
      delta_init <- if (object$include_outside_option) {
        object$coefficients[pm$asc]            # length J_inside (all ASCs free)
      } else {
        c(0, object$coefficients[pm$asc])      # length J_inside, baseline = 0
      }
    } else {
      delta_init <- rep(0, J)
    }
  }

  blp_contraction(
    delta = delta_init,
    target_shares = target_shares,
    X = d$X,
    beta = beta,
    alt_idx = d$alt_idx,
    M = d$M,
    weights = d$weights,
    include_outside_option = object$include_outside_option,
    tol = tol,
    max_iter = max_iter
  )
}

# --- elasticities: MXL -------------------------------------------------------

#' Elasticities for mixed logit model
#'
#' @param object A \code{choicer_mxl} object fitted with \code{keep_data = TRUE}.
#' @param elast_var Variable for elasticity computation: a column name (character)
#'   or 1-based index. Indexes into X columns for fixed coefficients, or W columns
#'   for random coefficients (when \code{is_random_coef = TRUE}).
#' @param is_random_coef Logical. \code{TRUE} if the variable has a random
#'   coefficient (is in W), \code{FALSE} if fixed (in X). Default \code{FALSE}.
#' @param ... Additional arguments (ignored).
#' @returns A J x J elasticity matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' elasticities(fit, "x1")
#' elasticities(fit, "w1", is_random_coef = TRUE)
#' }
#' @export
elasticities.choicer_mxl <- function(object, elast_var,
                                     is_random_coef = FALSE, ...) {
  if (is.null(object[["data"]])) {
    stop("elasticities() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]

  col_names <- if (is_random_coef) colnames(d$W) else colnames(d$X)
  idx <- resolve_var_index(elast_var, col_names)

  gp_el <- .mxl_gen_params(object$draws_info)

  mat <- mxl_elasticities_parallel(
    theta = object$coefficients,
    X = d$X,
    W = d$W,
    alt_idx = d$alt_idx,
    choice_idx = d$choice_idx,
    M = d$M,
    weights = d$weights,
    eta_draws = gp_el$eta_draws,
    rc_dist = object$rc_dist,
    elast_var_idx = idx,
    is_random_coef = is_random_coef,
    rc_correlation = object$rc_correlation,
    rc_mean = object$rc_mean,
    use_asc = object$use_asc,
    include_outside_option = object$include_outside_option,
    gen_seed = gp_el$gen_seed, gen_scramble = gp_el$gen_scramble, gen_S = gp_el$gen_S
  )

  label_matrix(mat, object$alt_mapping)
}

# --- diversion_ratios: MXL ---------------------------------------------------

#' Diversion ratios for mixed logit model
#'
#' Computes the attribute-based diversion ratio matrix. Entry (k, j) is the
#' fraction of demand lost by alternative j that is captured by alternative k
#' when a marginal change in alternative j's \code{wrt_var} attribute reduces
#' s_j.
#'
#' Unlike MNL, the MXL diversion ratio depends on which variable is perturbed:
#' the realised coefficient \eqn{\beta_{ik}^s} varies across individuals and
#' draws and does not cancel in the ratio. For a variable with a fixed
#' coefficient the result is independent of the variable (\eqn{\beta} cancels);
#' for a random-coefficient variable it is not.
#'
#' @param object A \code{choicer_mxl} object fitted with \code{keep_data = TRUE}.
#' @param wrt_var Variable used to perturb alternative j's utility: a column
#'   name (character) or 1-based index. Indexes into X columns for fixed
#'   coefficients, or W columns for random coefficients (when
#'   \code{is_random_coef = TRUE}).
#' @param is_random_coef Logical. \code{TRUE} if the variable has a random
#'   coefficient (is in W), \code{FALSE} if fixed (in X). Default \code{FALSE}.
#' @param ... Additional arguments (ignored).
#' @returns A J x J diversion ratio matrix with alternative labels.
#'   Cross-products are averaged across simulation draws inside the
#'   integration to avoid Jensen-style bias.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' diversion_ratios(fit, "x1")
#' diversion_ratios(fit, "w1", is_random_coef = TRUE)
#' }
#' @export
diversion_ratios.choicer_mxl <- function(object, wrt_var,
                                         is_random_coef = FALSE, ...) {
  if (is.null(object[["data"]])) {
    stop("diversion_ratios() requires stored data. Refit with keep_data = TRUE.")
  }
  if (is.null(object$draws_info)) {
    stop("diversion_ratios() requires draws_info from a fitted MXL model.")
  }
  d <- object[["data"]]

  col_names <- if (is_random_coef) colnames(d$W) else colnames(d$X)
  idx <- resolve_var_index(wrt_var, col_names)

  gp_dr <- .mxl_gen_params(object$draws_info)

  mat <- mxl_diversion_ratios_parallel(
    theta                  = object$coefficients,
    X                      = d$X,
    W                      = d$W,
    alt_idx                = d$alt_idx,
    M                      = d$M,
    weights                = d$weights,
    eta_draws              = gp_dr$eta_draws,
    rc_dist                = object$rc_dist,
    elast_var_idx          = idx,
    is_random_coef         = is_random_coef,
    rc_correlation         = object$rc_correlation,
    rc_mean                = object$rc_mean,
    use_asc                = object$use_asc,
    include_outside_option = object$include_outside_option,
    gen_seed = gp_dr$gen_seed, gen_scramble = gp_dr$gen_scramble, gen_S = gp_dr$gen_S
  )

  label_matrix(mat, object$alt_mapping)
}

# --- blp: MXL ----------------------------------------------------------------

#' BLP contraction mapping for mixed logit model
#'
#' @param object A \code{choicer_mxl} object fitted with \code{keep_data = TRUE}.
#' @param target_shares Numeric vector of target market shares.
#'   Length \code{J_inside} when no outside option, or \code{J_inside + 1}
#'   (with the outside option's share at index 1) when
#'   \code{include_outside_option = TRUE}.
#' @param delta_init Initial guess for delta (ASC) values. If \code{NULL},
#'   uses the estimated ASCs from the fitted model.
#' @param tol Convergence tolerance (default 1e-8).
#' @param max_iter Maximum iterations (default 1000).
#' @param ... Additional arguments (ignored).
#' @returns Converged delta (ASC) vector.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 3
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, `:=`(x1 = rnorm(.N), w1 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_mxlogit(
#'   data = dt, id_col = "id", alt_col = "alt", choice_col = "choice",
#'   covariate_cols = "x1", random_var_cols = "w1", S = 50L
#' )
#' blp(fit, target_shares = rep(1/J, J))
#' }
#' @export
blp.choicer_mxl <- function(object, target_shares, delta_init = NULL,
                            tol = 1e-8, max_iter = 1000, ...) {
  if (is.null(object[["data"]])) {
    stop("blp() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]
  pm <- object$param_map

  beta <- object$coefficients[pm$beta]
  mu <- if (!is.null(pm$mu)) object$coefficients[pm$mu] else rep(0, object$draws_info$K_w)
  L_params <- object$coefficients[pm$sigma]

  J <- nrow(object$alt_mapping)
  if (is.null(delta_init)) {
    if (!is.null(pm$asc)) {
      delta_init <- if (object$include_outside_option) {
        object$coefficients[pm$asc]            # length J_inside (all ASCs free)
      } else {
        c(0, object$coefficients[pm$asc])      # length J_inside, baseline = 0
      }
    } else {
      delta_init <- rep(0, J)
    }
  }

  gp_blp <- .mxl_gen_params(object$draws_info)

  mxl_blp_contraction(
    delta = delta_init,
    target_shares = target_shares,
    X = d$X,
    W = d$W,
    beta = beta,
    mu = mu,
    L_params = L_params,
    alt_idx = d$alt_idx,
    M = d$M,
    weights = d$weights,
    eta_draws = gp_blp$eta_draws,
    rc_dist = object$rc_dist,
    rc_correlation = object$rc_correlation,
    rc_mean = object$rc_mean,
    include_outside_option = object$include_outside_option,
    tol = tol,
    max_iter = max_iter,
    gen_seed = gp_blp$gen_seed,
    gen_scramble = gp_blp$gen_scramble,
    gen_S = gp_blp$gen_S
  )
}

# --- predict: NL -------------------------------------------------------------

#' Predict from a nested logit model
#'
#' Computes choice probabilities or aggregate market shares, either for the
#' data used at fit time (default) or for counterfactual `newdata`.
#'
#' @param object A choicer_nl object.
#' @param type One of "probabilities" (individual-level choice probabilities)
#'   or "shares" (aggregate market shares).
#' @param newdata Optional data for counterfactual prediction. Either:
#'   * a data.frame in the same long format used at fit time (one row per
#'     id-alternative pair, with the fit-time id, alternative, and covariate
#'     columns; a choice column is not required). Alternative labels must have
#'     been seen at fit time; per-id subsets of alternatives are allowed. The
#'     alternative-to-nest mapping always comes from the fitted object (it
#'     indexes the estimated `lambda` parameters), so a nest column in
#'     `newdata` is not required and is ignored if present.
#'   * a list with elements `X`, `alt_idx`, `M` (and optionally `weights`)
#'     matching the layout of `object$data` — the "modified design matrix"
#'     path for policy simulation. `alt_idx` must use the fit-time integer
#'     codes from `object$alt_mapping`.
#'
#'   When `NULL` (default), the data stored at fit time is used (requires
#'   `keep_data = TRUE`).
#' @param weights Optional numeric vector with one weight per choice situation,
#'   used for `type = "shares"` aggregation. For a data.frame `newdata`,
#'   supply one weight per id in order of first appearance in `newdata`
#'   (weights are realigned internally to the sorted row order). Defaults to
#'   equal weights. Ignored when `newdata` is `NULL` (the stored fit weights
#'   apply).
#' @param ... Additional arguments (ignored).
#' @returns For "probabilities": a list with `choice_prob` and `utility` vectors.
#'   For "shares": a named numeric vector of market shares per alternative.
#'   With a data.frame `newdata`, rows are ordered by id, then by fit-time
#'   alternative code (`alt_int` in `object$alt_mapping`).
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, nest := rep(c(1L, 1L, 2L, 2L), N)]
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(dt, "id", "alt", "choice", c("x1", "x2"), "nest")
#' predict(fit, type = "shares")
#' predict(fit, type = "probabilities")
#' }
#' @export
predict.choicer_nl <- function(object, type = c("probabilities", "shares"),
                               newdata = NULL, weights = NULL, ...) {
  type <- match.arg(type)

  if (is.null(newdata)) {
    if (is.null(object[["data"]])) {
      stop("Prediction requires stored data. Refit with keep_data = TRUE.")
    }
    d <- object[["data"]]
    nest_idx <- d$nest_idx
  } else {
    # The per-alternative (length J) nest mapping is authoritative from the
    # fit since it indexes the estimated lambda parameters. New fits store it
    # top-level; older fits only carry it inside $data (keep_data = TRUE).
    nest_idx <- object$nest_idx %||% object[["data"]]$nest_idx
    if (is.null(nest_idx)) {
      stop("Cannot resolve the alternative-to-nest mapping: this fit stores ",
           "no 'nest_idx'. Refit to enable newdata prediction.")
    }
    d <- resolve_predict_newdata(object, newdata, weights = weights)
  }

  theta <- object$coefficients

  if (type == "probabilities") {
    nl_predict(
      theta = theta,
      X = d$X,
      alt_idx = d$alt_idx,
      M = d$M,
      nest_idx = nest_idx,
      use_asc = object$use_asc,
      include_outside_option = object$include_outside_option
    )
  } else {
    nl_predict_shares(
      theta = theta,
      X = d$X,
      alt_idx = d$alt_idx,
      M = d$M,
      weights = d$weights,
      nest_idx = nest_idx,
      use_asc = object$use_asc,
      include_outside_option = object$include_outside_option
    )
  }
}

# --- elasticities: NL --------------------------------------------------------

#' Elasticities for nested logit model
#'
#' @param object A \code{choicer_nl} object fitted with \code{keep_data = TRUE}.
#' @param elast_var Variable for elasticity computation: a column name (character)
#'   or 1-based index into the design matrix X.
#' @param ... Additional arguments (ignored).
#' @returns A J x J elasticity matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, nest := rep(c(1L, 1L, 2L, 2L), N)]
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(dt, "id", "alt", "choice", c("x1", "x2"), "nest")
#' elasticities(fit, "x1")
#' }
#' @export
elasticities.choicer_nl <- function(object, elast_var, ...) {
  if (is.null(object[["data"]])) {
    stop("elasticities() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]
  idx <- resolve_var_index(elast_var, colnames(d$X))

  mat <- nl_elasticities_parallel(
    theta = object$coefficients,
    X = d$X,
    alt_idx = d$alt_idx,
    choice_idx = d$choice_idx,
    nest_idx = d$nest_idx,
    M = d$M,
    weights = d$weights,
    elast_var_idx = idx,
    use_asc = object$use_asc,
    include_outside_option = object$include_outside_option
  )

  label_matrix(mat, object$alt_mapping)
}

# --- diversion_ratios: NL ----------------------------------------------------

#' Diversion ratios for nested logit model
#'
#' @param object A \code{choicer_nl} object fitted with \code{keep_data = TRUE}.
#' @param ... Additional arguments (ignored).
#' @returns A J x J diversion ratio matrix with alternative labels.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, nest := rep(c(1L, 1L, 2L, 2L), N)]
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(dt, "id", "alt", "choice", c("x1", "x2"), "nest")
#' diversion_ratios(fit)
#' }
#' @export
diversion_ratios.choicer_nl <- function(object, ...) {
  if (is.null(object[["data"]])) {
    stop("diversion_ratios() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]

  mat <- nl_diversion_ratios_parallel(
    theta = object$coefficients,
    X = d$X,
    alt_idx = d$alt_idx,
    nest_idx = d$nest_idx,
    M = d$M,
    weights = d$weights,
    use_asc = object$use_asc,
    include_outside_option = object$include_outside_option
  )

  label_matrix(mat, object$alt_mapping)
}

# --- blp: NL -----------------------------------------------------------------

#' BLP contraction mapping for nested logit model
#'
#' @param object A \code{choicer_nl} object fitted with \code{keep_data = TRUE}.
#' @param target_shares Numeric vector of target market shares.
#'   Length \code{J_inside} when no outside option, or \code{J_inside + 1}
#'   (with the outside option's share at index 1) when
#'   \code{include_outside_option = TRUE}.
#' @param delta_init Initial guess for delta (ASC) values. If \code{NULL},
#'   uses the estimated ASCs from the fitted model.
#' @param damping Contraction damping factor in (0, 1] (default 1).
#' @param tol Convergence tolerance (default 1e-8).
#' @param max_iter Maximum iterations (default 1000).
#' @param ... Additional arguments (ignored).
#' @returns Converged delta (ASC) vector.
#' @examples
#' \donttest{
#' library(data.table)
#' set.seed(42)
#' N <- 50; J <- 4
#' dt <- data.table(id = rep(1:N, each = J), alt = rep(1:J, N))
#' dt[, nest := rep(c(1L, 1L, 2L, 2L), N)]
#' dt[, `:=`(x1 = rnorm(.N), x2 = rnorm(.N))]
#' dt[, choice := 0L]
#' dt[, choice := sample(c(1L, rep(0L, J - 1))), by = id]
#' fit <- run_nestlogit(dt, "id", "alt", "choice", c("x1", "x2"), "nest")
#' blp(fit, target_shares = rep(1/J, J))
#' }
#' @export
blp.choicer_nl <- function(object, target_shares, delta_init = NULL,
                           damping = 1, tol = 1e-8, max_iter = 1000, ...) {
  if (is.null(object[["data"]])) {
    stop("blp() requires stored data. Refit with keep_data = TRUE.")
  }
  d <- object[["data"]]
  pm <- object$param_map
  beta <- object$coefficients[pm$beta]

  # Build the full length-n_nests lambda vector expected by the kernel.
  # Singleton nests have lambda fixed to 1; non-singleton nests receive the
  # estimated lambdas, scattered by ascending nest index (matching the C++).
  n_nests <- max(d$nest_idx)
  lambda_full <- rep(1, n_nests)
  nest_counts <- tabulate(d$nest_idx, n_nests)
  non_singleton <- which(nest_counts > 1)
  lambda_full[non_singleton] <- object$coefficients[pm$lambda]

  J <- nrow(object$alt_mapping)
  if (is.null(delta_init)) {
    if (!is.null(pm$asc)) {
      delta_init <- if (object$include_outside_option) {
        object$coefficients[pm$asc]            # length J_inside (all ASCs free)
      } else {
        c(0, object$coefficients[pm$asc])      # length J_inside, baseline = 0
      }
    } else {
      delta_init <- rep(0, J)
    }
  }

  nl_blp_contraction(
    delta = delta_init,
    target_shares = target_shares,
    X = d$X,
    beta = beta,
    lambda = lambda_full,
    alt_idx = d$alt_idx,
    nest_idx = d$nest_idx,
    M = d$M,
    weights = d$weights,
    include_outside_option = object$include_outside_option,
    damping = damping,
    tol = tol,
    max_iter = max_iter
  )
}

# --- choicer_hb (hierarchical Bayes) methods ---------------------------------
# Shared by choicer_hmnl and choicer_hmnp: posterior-draws objects with a
# b-coefficient block, a delta/xi quality ladder, and hierarchy summaries.
# Written once on the choicer_hb parent (not choicer_fit -- no loglik /
# Hessian machinery applies).

#' Print a hierarchical Bayes fit
#'
#' @param x A `choicer_hmnl` or `choicer_hmnp` object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @examples
#' \donttest{
#' sim <- simulate_hmnl_data(N = 50, T = 2, J = 3, seed = 42)
#' fit <- suppressWarnings(run_hmnlogit(sim$data, "task", "alt", "choice", c("x1", "x2"),
#'                     person_col = "pid",
#'                     mcmc = list(R = 300, burn = 100)))
#' print(fit)
#' }
#' @export
print.choicer_hb <- function(x, ...) {
  cat(model_display_name(x$model), "model\n")
  cat("Posterior-mean population coefficients (b):\n")
  print(round(x$coefficients, 4))
  cat("\nAlternatives:", x$J, " Respondents:", x$n_persons,
      " Choice situations:", x$nobs, "\n")
  cat("Draws kept:", x$mcmc$R_keep, " (R =", x$mcmc$R, ", burn =",
      x$mcmc$burn, ", thin =", x$mcmc$thin, ")\n")
  invisible(x)
}

#' Summarize a hierarchical Bayes fit
#'
#' Posterior summaries (mean, SD, equal-tailed credible interval) for the
#' population coefficients \eqn{b}, the mean-function coefficients
#' \eqn{\theta}, the alternative-effect variance \eqn{\sigma_d^2} (and, for
#' the HMNP, the raw shock variance trace), plus the \eqn{\delta_j} /
#' \eqn{\xi_j} quality ladder, acceptance diagnostics, and a consolidated
#' convergence-diagnostic table (rank-normalized R-hat, ESS bulk/tail, MCSE)
#' built from all retained chains (see [rhat()], [ess()], [mcse()]).
#'
#' @param object A `choicer_hmnl` or `choicer_hmnp` object.
#' @param prob Probability mass of the equal-tailed credible interval
#'   (default 0.95).
#' @param ... Additional arguments (ignored).
#' @returns A `summary.choicer_hb` object.
#' @examples
#' \donttest{
#' sim <- simulate_hmnl_data(N = 50, T = 2, J = 3, seed = 42)
#' fit <- suppressWarnings(run_hmnlogit(sim$data, "task", "alt", "choice", c("x1", "x2"),
#'                     person_col = "pid",
#'                     mcmc = list(R = 300, burn = 100)))
#' summary(fit)
#' }
#' @export
summary.choicer_hb <- function(object, prob = 0.95, ...) {
  if (!is.numeric(prob) || length(prob) != 1L || prob <= 0 || prob >= 1) {
    stop("prob must be a single number in (0, 1).")
  }
  structure(
    list(
      model = object$model,
      coefficients = build_bayes_coef_table(object$draws$b, prob),
      theta_table = build_bayes_coef_table(object$draws$theta, prob),
      sigma_d2_table = build_bayes_coef_table(
        matrix(object$draws$sigma_d2, ncol = 1,
               dimnames = list(NULL, "sigma_d^2")), prob),
      sigma2_table = if (!is.null(object$draws$sigma2)) {
        build_bayes_coef_table(
          matrix(object$draws$sigma2, ncol = 1,
                 dimnames = list(NULL, "sigma^2 (raw)")), prob)
      },
      delta = object$delta,
      xi = object$xi,
      W_mean = object$W_mean,
      accept = object$accept,
      rhat = object$rhat,
      diagnostic_table = .hb_diagnostic_table(object),
      prob = prob,
      nobs = object$nobs,
      n_persons = object$n_persons,
      J = object$J,
      mcmc = object$mcmc,
      elapsed_time = object$sampler$elapsed_time
    ),
    class = "summary.choicer_hb"
  )
}

#' Print the summary of a hierarchical Bayes fit
#'
#' @param x A `summary.choicer_hb` object.
#' @param ... Additional arguments (ignored).
#' @returns The object invisibly.
#' @export
print.summary.choicer_hb <- function(x, ...) {
  cat(model_display_name(x$model), "model\n\n")
  cat("Population coefficients b (posterior):\n")
  print_bayes_coef_table(x$coefficients, x$prob)
  cat("\nDelta mean function theta (posterior):\n")
  print_bayes_coef_table(x$theta_table, x$prob)
  cat("\nAlternative-effect variance (posterior):\n")
  print_bayes_coef_table(x$sigma_d2_table, x$prob)
  if (!is.null(x$sigma2_table)) {
    cat("\nRaw shock variance (non-identified chain):\n")
    print_bayes_coef_table(x$sigma2_table, x$prob)
  }
  cat("\nQuality ladder (delta = mean utility vs the outside option;",
      "xi = delta - z'theta):\n")
  ladder <- data.frame(
    alternative = x$delta$alternative,
    delta_mean = round(x$delta$mean, 4),
    delta_sd = round(x$delta$sd, 4),
    xi_mean = round(x$xi$mean, 4),
    xi_sd = round(x$xi$sd, 4)
  )
  print(ladder, row.names = FALSE)
  cat(sprintf("\nConvergence diagnostics (%d chain%s, %d draws each)\n",
              x$diagnostic_table$chains,
              if (x$diagnostic_table$chains == 1) "" else "s",
              x$diagnostic_table$R_keep))
  .print_hb_diagnostic_table(x$diagnostic_table)
  if (!is.null(x$accept) && !is.null(x$accept$mean_beta)) {
    cat(sprintf("Acceptance: beta %.2f, delta %.2f\n",
                x$accept$mean_beta, x$accept$mean_delta))
  } else {
    cat("Acceptance: conjugate \u2014 no acceptance step\n")
  }
  cat("\nRespondents:", x$n_persons, " Choice situations:", x$nobs,
      " Alternatives:", x$J, "\n")
  cat("Draws kept:", x$mcmc$R_keep, " Chains:", x$mcmc$chains, "\n")
  if (!is.null(x$elapsed_time)) cat("MCMC run time", x$elapsed_time, "\n")
  invisible(x)
}

#' Extract posterior means from a hierarchical Bayes fit
#'
#' @param object A `choicer_hmnl` or `choicer_hmnp` object.
#' @param component Which block to return: `"beta"` (population means b,
#'   default), `"theta"` (delta mean function), `"delta"` (alternative
#'   effects), or `"xi"` (unobserved quality, delta - z'theta).
#' @param ... Additional arguments (ignored).
#' @returns Named numeric vector of posterior means.
#' @examples
#' \donttest{
#' sim <- simulate_hmnl_data(N = 50, T = 2, J = 3, seed = 42)
#' fit <- suppressWarnings(run_hmnlogit(sim$data, "task", "alt", "choice", c("x1", "x2"),
#'                     person_col = "pid",
#'                     mcmc = list(R = 300, burn = 100)))
#' coef(fit)
#' coef(fit, component = "delta")
#' }
#' @export
coef.choicer_hb <- function(object,
                            component = c("beta", "theta", "delta", "xi"),
                            ...) {
  component <- match.arg(component)
  switch(component,
    beta = object$coefficients,
    theta = stats::setNames(colMeans(object$draws$theta),
                            colnames(object$draws$theta)),
    delta = stats::setNames(object$delta$mean, object$delta$alternative),
    xi = stats::setNames(object$xi$mean, object$xi$alternative)
  )
}

#' Posterior covariance of the population coefficients
#'
#' @param object A `choicer_hmnl` or `choicer_hmnp` object.
#' @param ... Additional arguments (ignored).
#' @returns K x K posterior covariance matrix of the b draws.
#' @export
vcov.choicer_hb <- function(object, ...) {
  object$vcov
}

#' Number of choice situations behind a hierarchical Bayes fit
#'
#' @param object A `choicer_hmnl` or `choicer_hmnp` object.
#' @param ... Additional arguments (ignored).
#' @returns Integer count of choice situations (tasks).
#' @export
nobs.choicer_hb <- function(object, ...) {
  object$nobs
}

Try the choicer package in your browser

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

choicer documentation built on Sept. 5, 2026, 1:07 a.m.