R/simulation.R

Defines functions `%||%` validate as.data.frame.biogeme_simulation print.biogeme_simulation evaluate_biogeme_expression_c simulate_single_formula bayesian_posterior_mean_by_observation simulate_bayesian biogeme_confidence_intervals predict.biogeme_bayesian_fit predict.biogeme_fit predict_biogeme_result biogeme_prediction_expressions biogeme_scenario_database simulate.biogeme_model simulate

Documented in bayesian_posterior_mean_by_observation biogeme_confidence_intervals evaluate_biogeme_expression_c predict.biogeme_fit simulate simulate_bayesian simulate_single_formula validate

#' Simulate named native Biogeme expressions at fixed estimates
#'
#' @param model A `biogeme_model`.
#' @param expressions Optional named list of expressions. When omitted, the
#'   model's `simulations` list is used.
#' @param beta A `biogeme_fit` or named numeric parameter vector.
#' @param control Optional [biogeme_control()] object.
#' @param database Optional scenario database. It may be a
#'   `biogeme_database` or a numeric data frame. A data frame inherits the
#'   model database's derived-variable, filter, and panel metadata.
#' @return A `biogeme_simulation` object containing a data frame of values.
#' @details
#' `expressions` is a named list of complete symbolic expressions. When it is
#' omitted, the model's `simulations` list is used. Native Biogeme evaluates
#' these expressions at the supplied parameter values and returns ordinary R
#' data frames.
#' @examples
#' \dontrun{
#' simulation_model <- biogeme_model(
#'   database,
#'   formula = choice_log_probability,
#'   simulations = list(probability = choice_probability)
#' )
#' simulated <- simulate(
#'   simulation_model,
#'   beta = fit,
#'   control = biogeme_control(output_directory = tempfile("rbiogeme-sim-"))
#' )
#' as.data.frame(simulated)
#' }
#' @export
simulate <- function(model, expressions = NULL, beta, control = NULL, database = NULL) {
  UseMethod("simulate")
}

#' @export
simulate.biogeme_model <- function(
    model,
    expressions = NULL,
    beta,
    control = NULL,
    database = NULL
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  if (inherits(beta, "biogeme_fit") || inherits(beta, "biogeme_bayesian_fit")) {
    beta <- stats::coef(beta)
  }
  if (!is.numeric(beta) || is.null(names(beta)) || anyNA(names(beta)) ||
      any(!nzchar(names(beta))) || anyNA(beta) || any(!is.finite(beta))) {
    stop("beta must be a named finite numeric vector or biogeme_fit.", call. = FALSE)
  }
  if (is.null(expressions)) expressions <- model$simulations
  if (is.null(expressions)) {
    if (!is.null(model$probability)) expressions <- list(probability = model$probability)
    else stop("Provide expressions or define model$simulations.", call. = FALSE)
  }
  if (!is.list(expressions) || is.null(names(expressions)) ||
      anyNA(names(expressions)) || any(!nzchar(names(expressions)))) {
    stop("expressions must be a named list of Biogeme expressions.", call. = FALSE)
  }
  expressions <- lapply(expressions, as_biogeme_expression)
  simulation_model <- model
  simulation_model$simulations <- expressions
  if (is.null(simulation_model$formula) && is.null(simulation_model$log_likelihood)) {
    simulation_model$formula <- NULL
  }
  controls <- control %||% model$control %||% list()
  if (!is.list(controls)) stop("control must be a named list.", call. = FALSE)
  database_override <- biogeme_scenario_database(model, database)
  compiled <- biogeme_compile_model(simulation_model, database_override = database_override)
  raw <- tryCatch(
    biogeme_bridge()$simulate_biogeme(
      compiled_model = compiled,
      beta_values = reticulate::r_to_py(as.list(beta)),
      model_name = if (!is.null(controls$model_name)) controls$model_name else "rbiogeme_simulation",
      controls = if (length(controls) == 0L) NULL else reticulate::r_to_py(controls)
    ),
    error = function(error) {
      biogeme_rethrow(
        error,
        class = "biogeme_estimation_error",
        operation = "Biogeme simulation",
        suggestion = "check the simulation expressions, parameter names, and draw controls"
      )
    }
  )
  result <- reticulate::py_to_r(raw)
  values <- result$values
  if (is.data.frame(values)) {
    values <- as.data.frame(values, check.names = FALSE)
  } else if (is.list(values)) {
    values <- as.data.frame(values, check.names = FALSE)
  } else {
    values <- data.frame(value = as.numeric(values))
  }
  structure(
    list(
      values = values,
      expressions = expressions,
      beta = beta,
      draw_types = result$draw_types,
      number_of_draws = result$number_of_draws,
      model = model
    ),
    class = "biogeme_simulation"
  )
}

biogeme_scenario_database <- function(model, database) {
  if (is.null(database)) return(NULL)
  validate_biogeme_database(model$database)
  if (inherits(database, "biogeme_database")) {
    validate_biogeme_database(database)
    if (!is.null(model$database$panel_id) && is.null(database$panel_id) &&
        biogeme_database_has_column(database, model$database$panel_id)) {
      database <- biogeme_database_panel(database, model$database$panel_id)
    }
    return(database)
  }
  if (!is.data.frame(database)) {
    stop("database must be a biogeme_database or numeric data.frame.", call. = FALSE)
  }
  scenario <- biogeme_database(model$database$name, database)
  if (!is.null(model$database$panel_id)) {
    scenario <- biogeme_database_panel(scenario, model$database$panel_id)
  }
  scenario$derived_variables <- model$database$derived_variables
  scenario$filters <- model$database$filters
  scenario$materialized <- length(scenario$derived_variables) == 0L &&
    length(scenario$filters) == 0L
  scenario
}

biogeme_prediction_expressions <- function(model) {
  if (!is.null(model$simulations)) return(model$simulations)
  if (!is.null(model$probability)) return(list(probability = model$probability))
  if (inherits(model, "biogeme_logit_model")) {
    return(stats::setNames(
      lapply(unname(model$alternative_codes), function(code) {
        logit_probability(
          utilities = model$utilities,
          availability = model$availability,
          alternative = code,
          alternative_codes = model$alternative_codes
        )
      }),
      names(model$utilities)
    ))
  }
  if (inherits(model, "biogeme_nested_logit_model")) {
    if (!is.null(model$scale_parameter)) {
      stop(
        "Default prediction is unavailable for a nested model with a custom scale parameter; provide expressions explicitly.",
        call. = FALSE
      )
    }
    return(stats::setNames(
      lapply(unname(model$alternative_codes), function(code) {
        nested_probability(
          utilities = model$utilities,
          availability = model$availability,
          nests = model$nests,
          alternative = code,
          alternative_codes = model$alternative_codes
        )
      }),
      names(model$utilities)
    ))
  }
  if (inherits(model, "biogeme_cross_nested_logit_model")) {
    if (!is.null(model$scale_parameter)) {
      stop(
        "Default prediction is unavailable for a cross-nested model with a custom scale parameter; provide expressions explicitly.",
        call. = FALSE
      )
    }
    return(stats::setNames(
      lapply(unname(model$alternative_codes), function(code) {
        cross_nested_probability(
          utilities = model$utilities,
          availability = model$availability,
          nests = model$nests,
          alternative = code,
          alternative_codes = model$alternative_codes
        )
      }),
      names(model$utilities)
    ))
  }
  stop(
    "Provide expressions or define model$probability/model$simulations for prediction.",
    call. = FALSE
  )
}

predict_biogeme_result <- function(object, newdata = NULL, expressions = NULL, control = NULL, ...) {
  if (!inherits(object, c("biogeme_fit", "biogeme_bayesian_fit"))) {
    stop("object must be a Biogeme estimation result.", call. = FALSE)
  }
  if (!inherits(object$model, "biogeme_model")) {
    stop("The estimation result does not contain its Biogeme model.", call. = FALSE)
  }
  extra <- list(...)
  if (length(extra) > 0L) {
    stop("Unsupported prediction arguments: ", paste(names(extra), collapse = ", "), call. = FALSE)
  }
  if (is.null(expressions)) expressions <- biogeme_prediction_expressions(object$model)
  simulation <- simulate(
    object$model,
    expressions = expressions,
    beta = object,
    control = control,
    database = newdata
  )
  as.data.frame(simulation, check.names = FALSE)
}

#' Predict native probabilities or simulation expressions
#'
#' @param object A `biogeme_fit` or `biogeme_bayesian_fit` object.
#' @param newdata Optional numeric data frame or `biogeme_database` used for a
#'   scenario prediction. When a data frame is supplied, the model database's
#'   derived-variable, filter, and panel metadata are retained.
#' @param expressions Optional named list of complete symbolic expressions.
#'   For logit, nested-logit, and cross-nested-logit fits, omission produces one
#'   native probability column per alternative. For a generic model, the model's
#'   probability or simulation expressions are used.
#' @param control Optional [biogeme_control()] object for native simulation.
#' @param ... Reserved for future prediction options; unsupported arguments are
#'   rejected explicitly.
#' @return A data frame containing values calculated by native Biogeme.
#' @details
#' `predict()` is a convenience wrapper around [simulate()]. It does not
#' calculate probabilities in R and does not refit the model. Use
#' [simulate()] directly when several named quantities, posterior draws, or a
#' custom simulation workflow are needed.
#' @examples
#' database <- biogeme_database(
#'   "demo",
#'   data.frame(choice = c(1, 2), x = c(1, 2))
#' )
#' model <- logit_model(
#'   database,
#'   choice = "choice",
#'   utilities = list(`1` = 0, `2` = biogeme_beta("b") * variable("x"))
#' )
#' # After fitting: predict(fit) or predict(fit, newdata = data.frame(...))
#' @export
predict.biogeme_fit <- function(object, newdata = NULL, expressions = NULL, control = NULL, ...) {
  predict_biogeme_result(object, newdata, expressions, control, ...)
}

#' @export
predict.biogeme_bayesian_fit <- function(object, newdata = NULL, expressions = NULL, control = NULL, ...) {
  predict_biogeme_result(object, newdata, expressions, control, ...)
}

#' Calculate native simulation confidence intervals
#'
#' Native Biogeme performs the repeated simulation and quantile calculation.
#' R supplies ordinary parameter mappings and receives two ordinary data
#' frames; no native simulation or result objects are exposed.
#'
#' @param model A `biogeme_model`.
#' @param beta_values A non-empty list of named finite numeric parameter
#'   vectors, typically obtained from a fitted result's bootstrap field.
#' @param expressions Optional named list of expressions. When omitted, the
#'   model's `simulations` list is used.
#' @param interval_size Confidence interval size, strictly between zero and
#'   one. The native default is 0.9.
#' @param control Optional [biogeme_control()] object.
#' @return A list with `left` and `right` data frames and native metadata.
#' @export
biogeme_confidence_intervals <- function(
    model,
    beta_values,
    expressions = NULL,
    interval_size = 0.9,
    control = NULL
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a biogeme_model.", call. = FALSE)
  }
  if (!is.list(beta_values) || length(beta_values) == 0L) {
    stop("beta_values must be a non-empty list of named numeric vectors.", call. = FALSE)
  }
  beta_values <- lapply(seq_along(beta_values), function(index) {
    draw <- beta_values[[index]]
    if (!is.numeric(draw) || is.null(names(draw)) || length(draw) == 0L ||
        anyNA(names(draw)) || any(!nzchar(names(draw))) ||
        anyDuplicated(names(draw)) || anyNA(draw) || any(!is.finite(draw))) {
      stop(
        "Each beta_values entry must be a non-empty named finite numeric vector.",
        call. = FALSE
      )
    }
    as.list(setNames(as.numeric(draw), names(draw)))
  })
  if (!is.numeric(interval_size) || length(interval_size) != 1L ||
      is.na(interval_size) || !is.finite(interval_size) ||
      interval_size <= 0 || interval_size >= 1) {
    stop("interval_size must be one numeric value strictly between zero and one.", call. = FALSE)
  }
  if (!is.null(expressions)) {
    if (!is.list(expressions) || is.null(names(expressions)) ||
        length(expressions) == 0L || anyNA(names(expressions)) ||
        any(!nzchar(names(expressions))) || anyDuplicated(names(expressions))) {
      stop("expressions must be a non-empty named list of expressions.", call. = FALSE)
    }
    model <- model
    model$simulations <- lapply(expressions, as_biogeme_expression)
  }
  if (is.null(model$simulations) || length(model$simulations) == 0L) {
    stop("Provide expressions or define model$simulations.", call. = FALSE)
  }
  controls <- control %||% model$control %||% list()
  if (!is.list(controls)) stop("control must be a named list.", call. = FALSE)
  compiled <- biogeme_compile_model(model)
  raw <- tryCatch(
    biogeme_bridge()$confidence_intervals_biogeme(
      compiled_model = compiled,
      beta_values = reticulate::r_to_py(beta_values),
      interval_size = as.numeric(interval_size),
      model_name = if (!is.null(controls$model_name)) {
        controls$model_name
      } else {
        "rbiogeme_confidence_intervals"
      },
      controls = if (length(controls) == 0L) NULL else reticulate::r_to_py(controls)
    ),
    error = function(error) {
      biogeme_rethrow(
        error,
        class = "biogeme_evaluation_error",
        operation = "native Biogeme confidence intervals",
        suggestion = "check the simulation expressions, parameter names, and interval size"
      )
    }
  )
  result <- reticulate::py_to_r(raw)
  as_table <- function(value, side) {
    if (is.data.frame(value)) return(as.data.frame(value, check.names = FALSE))
    if (is.list(value)) return(as.data.frame(value, check.names = FALSE))
    stop("Biogeme returned an invalid ", side, " confidence-interval table.", call. = FALSE)
  }
  list(
    left = as_table(result$left, "left"),
    right = as_table(result$right, "right"),
    interval_size = as.numeric(result$interval_size),
    number_of_parameter_draws = as.integer(result$number_of_parameter_draws)
  )
}

#' Simulate named formulas over native Bayesian posterior draws
#'
#' The NetCDF file is loaded by native Biogeme and each selected posterior draw
#' is evaluated by native Biogeme. R receives only the native summary table;
#' PyMC and ArviZ objects are never exposed to ordinary R code.
#'
#' @param model A `biogeme_model`.
#' @param bayesian_results A `biogeme_bayesian_fit` or path to a native NetCDF
#'   Bayesian result file.
#' @param expressions Optional named list of expressions. When omitted, the
#'   model's `simulations` list is used.
#' @param percentage_of_draws_to_use Percentage of posterior draws passed to
#'   native Biogeme for simulation.
#' @param lower_quantile Lower posterior-simulation quantile.
#' @param upper_quantile Upper posterior-simulation quantile.
#' @param control Optional [biogeme_control()] object.
#' @return A `biogeme_simulation` object containing native posterior summaries.
#' @export
simulate_bayesian <- function(
    model,
    bayesian_results,
    expressions = NULL,
    percentage_of_draws_to_use = 10,
    lower_quantile = 0.025,
    upper_quantile = 0.975,
    control = NULL
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  if (inherits(bayesian_results, "biogeme_bayesian_fit")) {
    bayesian_results <- bayesian_results$netcdf_file
  }
  if (!is.character(bayesian_results) || length(bayesian_results) != 1L ||
      is.na(bayesian_results) || !nzchar(bayesian_results)) {
    stop("bayesian_results must be a biogeme_bayesian_fit or one NetCDF path.", call. = FALSE)
  }
  bayesian_results <- path.expand(bayesian_results)
  if (!file.exists(bayesian_results) || dir.exists(bayesian_results)) {
    stop("The Bayesian NetCDF result file does not exist: ", bayesian_results, call. = FALSE)
  }
  validate_probability <- function(value, name) {
    if (!is.numeric(value) || length(value) != 1L || is.na(value) ||
        !is.finite(value) || value <= 0 || value >= 1) {
      stop(name, " must be one numeric value strictly between zero and one.", call. = FALSE)
    }
    as.numeric(value)
  }
  if (!is.numeric(percentage_of_draws_to_use) || length(percentage_of_draws_to_use) != 1L ||
      is.na(percentage_of_draws_to_use) || !is.finite(percentage_of_draws_to_use) ||
      percentage_of_draws_to_use <= 0) {
    stop("percentage_of_draws_to_use must be one positive numeric value.", call. = FALSE)
  }
  lower_quantile <- validate_probability(lower_quantile, "lower_quantile")
  upper_quantile <- validate_probability(upper_quantile, "upper_quantile")
  if (lower_quantile >= upper_quantile) {
    stop("lower_quantile must be smaller than upper_quantile.", call. = FALSE)
  }
  if (!is.null(expressions)) {
    if (!is.list(expressions) || is.null(names(expressions)) ||
        anyNA(names(expressions)) || any(!nzchar(names(expressions)))) {
      stop("expressions must be a named list of Biogeme expressions.", call. = FALSE)
    }
    model <- model
    model$simulations <- lapply(expressions, as_biogeme_expression)
  }
  expressions <- model$simulations
  if (is.null(expressions) || length(expressions) == 0L) {
    stop("Provide expressions or define model$simulations.", call. = FALSE)
  }
  controls <- control %||% model$control %||% list()
  if (!is.list(controls)) stop("control must be a named list.", call. = FALSE)
  model_name <- if (!is.null(controls$model_name)) {
    controls$model_name
  } else {
    "rbiogeme_bayesian_simulation"
  }
  model_name <- validate_name(model_name, "model_name")
  controls <- validate_estimation_controls(controls)
  raw <- biogeme_simulate_bayesian_model(
    model = model,
    bayesian_results_file = bayesian_results,
    model_name = model_name,
    percentage_of_draws_to_use = percentage_of_draws_to_use,
    lower_quantile = lower_quantile,
    upper_quantile = upper_quantile,
    controls = controls
  )
  values <- raw$values
  if (is.data.frame(values)) {
    values <- as.data.frame(values, check.names = FALSE)
  } else if (is.list(values)) {
    values <- as.data.frame(values, check.names = FALSE)
  } else {
    values <- data.frame(value = as.numeric(values))
  }
  posterior_means <- as.numeric(unlist(raw$posterior_means, use.names = FALSE))
  names(posterior_means) <- names(raw$posterior_means)
  structure(
    list(
      values = values,
      expressions = expressions,
      beta = posterior_means,
      posterior_means = posterior_means,
      posterior_draws = as.integer(raw$posterior_draws),
      chains = as.integer(raw$chains),
      draws = as.integer(raw$draws),
      bayesian_results_file = as.character(raw$bayesian_results_file),
      model = model
    ),
    class = "biogeme_simulation"
  )
}

#' Retrieve native posterior means by observation
#'
#' Loads a Bayesian NetCDF result through Biogeme's public
#' `BayesianResults.from_netcdf()` API and delegates the observation-level
#' reduction to `posterior_mean_by_observation()`. R receives only the native
#' result table; posterior draws and ArviZ objects remain outside R.
#'
#' @param bayesian_results A `biogeme_bayesian_fit` or path to a native NetCDF
#'   Bayesian result file.
#' @param variable_name Name of a stored posterior variable with exactly one
#'   observation dimension in addition to chain and draw.
#' @param control Optional [biogeme_control()] object. The likelihood, WAIC,
#'   and LOO calculations are disabled by default because this operation only
#'   needs the stored posterior variable.
#' @return A data frame indexed by the native observation coordinate, with one
#'   column named `variable_name`.
#' @export
bayesian_posterior_mean_by_observation <- function(
    bayesian_results,
    variable_name,
    control = NULL
) {
  if (inherits(bayesian_results, "biogeme_bayesian_fit")) {
    bayesian_results <- bayesian_results$netcdf_file
  }
  if (!is.character(bayesian_results) || length(bayesian_results) != 1L ||
      is.na(bayesian_results) || !nzchar(bayesian_results)) {
    stop("bayesian_results must be a biogeme_bayesian_fit or one NetCDF path.", call. = FALSE)
  }
  bayesian_results <- path.expand(bayesian_results)
  if (!file.exists(bayesian_results) || dir.exists(bayesian_results)) {
    stop("The Bayesian NetCDF result file does not exist: ", bayesian_results, call. = FALSE)
  }
  variable_name <- validate_name(variable_name, "variable_name")
  controls <- if (is.null(control)) list() else control
  if (!is.list(controls)) stop("control must be a named list.", call. = FALSE)
  if (length(controls) > 0L) controls <- validate_estimation_controls(controls)
  raw <- biogeme_posterior_mean_by_observation(
    bayesian_results_file = bayesian_results,
    variable_name = variable_name,
    controls = controls
  )
  table <- raw$table
  if (!is.list(table) || is.null(table$index) || is.null(table$columns)) {
    stop("Biogeme returned an invalid posterior observation table.", call. = FALSE)
  }
  index <- as.character(table$index)
  columns <- as.character(table$columns)
  rows <- if (is.null(table$data)) list() else table$data
  if (length(columns) != 1L || !identical(columns, variable_name)) {
    stop("Biogeme returned an unexpected posterior observation table.", call. = FALSE)
  }
  values <- if (length(rows) == 0L) {
    numeric()
  } else {
    vapply(rows, function(row) {
      value <- unlist(row, use.names = FALSE)
      if (length(value) != 1L) {
        stop("Biogeme returned a non-scalar observation value.", call. = FALSE)
      }
      as.numeric(value[[1L]])
    }, numeric(1))
  }
  result <- data.frame(values, check.names = FALSE, stringsAsFactors = FALSE)
  names(result) <- variable_name
  rownames(result) <- index
  result
}

#' Evaluate one native Biogeme formula as an aggregated scalar
#'
#' This operation delegates to Biogeme's public
#' `calculate_single_formula_from_expression` evaluator. Unlike [simulate()],
#' it returns one native aggregated value, which is useful for panel
#' likelihood formulas.
#'
#' @param model A `biogeme_model`.
#' @param expression A Biogeme expression to evaluate.
#' @param beta A `biogeme_fit` or named numeric parameter vector.
#' @param number_of_draws Positive number of native Monte Carlo draws.
#' @param seed Optional temporary native draw seed.
#' @param numerically_safe Whether to request native numerically safe formulas.
#' @param use_jit Whether to use native JAX just-in-time compilation.
#' @return One numeric scalar returned by native Biogeme.
#' @export
simulate_single_formula <- function(
    model,
    expression,
    beta,
    number_of_draws,
    seed = NULL,
    numerically_safe = FALSE,
    use_jit = TRUE
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  expression <- as_biogeme_expression(expression)
  if (inherits(beta, "biogeme_fit")) beta <- stats::coef(beta)
  if (!is.numeric(beta) || is.null(names(beta)) || anyNA(names(beta)) ||
      any(!nzchar(names(beta))) || anyNA(beta) || any(!is.finite(beta))) {
    stop("beta must be a named finite numeric vector or biogeme_fit.", call. = FALSE)
  }
  if (!is.numeric(number_of_draws) || length(number_of_draws) != 1L ||
      is.na(number_of_draws) || !is.finite(number_of_draws) ||
      number_of_draws < 1 || number_of_draws != floor(number_of_draws)) {
    stop("number_of_draws must be one positive integer.", call. = FALSE)
  }
  if (!is.null(seed) && (!is.numeric(seed) || length(seed) != 1L ||
      is.na(seed) || !is.finite(seed) || seed != floor(seed))) {
    stop("seed must be NULL or one finite integer.", call. = FALSE)
  }
  if (!is.logical(numerically_safe) || length(numerically_safe) != 1L ||
      is.na(numerically_safe) || !is.logical(use_jit) || length(use_jit) != 1L ||
      is.na(use_jit)) {
    stop("numerically_safe and use_jit must be one non-missing logical value each.", call. = FALSE)
  }

  single_formula_model <- model
  single_formula_model$simulations <- list(single_formula = expression)
  compiled <- biogeme_compile_model(single_formula_model)
  raw <- tryCatch(
    biogeme_bridge()$calculate_single_formula_biogeme(
      compiled_model = compiled,
      formula_name = "single_formula",
      beta_values = reticulate::r_to_py(as.list(beta)),
      number_of_draws = as.integer(number_of_draws),
      seed = if (is.null(seed)) NULL else as.integer(seed),
      numerically_safe = numerically_safe,
      use_jit = use_jit
    ),
    error = function(error) {
      biogeme_rethrow(
        error,
        class = "biogeme_estimation_error",
        operation = "single-formula Biogeme evaluation",
        suggestion = "check the expression, parameter names, panel database, and draw count"
      )
    }
  )
  as.numeric(reticulate::py_to_r(raw)$value)
}

#' Evaluate one expression with native row-wise or aggregated JAX calculation
#'
#' This operation delegates to Biogeme's public `get_value_c` calculator. The
#' expression is compiled before the native evaluator runs; R receives only a
#' numeric vector or scalar.
#'
#' @param model A `biogeme_model` supplying the database.
#' @param expression A Biogeme expression to evaluate.
#' @param beta A `biogeme_fit` or named numeric parameter vector.
#' @param aggregation Whether to return one native aggregated scalar instead
#'   of one value per observation.
#' @param number_of_draws Positive number of native Monte Carlo draws.
#' @param numerically_safe Whether to request native numerically safe formulas.
#' @param use_jit Whether to use native JAX just-in-time compilation.
#' @return A numeric vector, or one numeric scalar when `aggregation = TRUE`.
#' @export
evaluate_biogeme_expression_c <- function(
    model,
    expression,
    beta,
    aggregation = FALSE,
    number_of_draws = 1000L,
    numerically_safe = FALSE,
    use_jit = TRUE
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  expression <- as_biogeme_expression(expression)
  if (inherits(beta, "biogeme_fit") || inherits(beta, "biogeme_bayesian_fit")) {
    beta <- stats::coef(beta)
  }
  if (!is.numeric(beta) || is.null(names(beta)) || anyNA(names(beta)) ||
      any(!nzchar(names(beta))) || anyDuplicated(names(beta)) || anyNA(beta) ||
      any(!is.finite(beta))) {
    stop("beta must be a named finite numeric vector or biogeme_fit.", call. = FALSE)
  }
  if (!is.logical(aggregation) || length(aggregation) != 1L || is.na(aggregation)) {
    stop("aggregation must be one non-missing logical value.", call. = FALSE)
  }
  if (!is.numeric(number_of_draws) || length(number_of_draws) != 1L ||
      is.na(number_of_draws) || !is.finite(number_of_draws) ||
      number_of_draws < 1 || number_of_draws != floor(number_of_draws)) {
    stop("number_of_draws must be one positive integer.", call. = FALSE)
  }
  if (!is.logical(numerically_safe) || length(numerically_safe) != 1L ||
      is.na(numerically_safe) || !is.logical(use_jit) || length(use_jit) != 1L ||
      is.na(use_jit)) {
    stop("numerically_safe and use_jit must be one non-missing logical value each.", call. = FALSE)
  }

  evaluation_model <- model
  evaluation_model$simulations <- list(rbiogeme_expression_c = expression)
  compiled <- biogeme_compile_model(evaluation_model)
  raw <- tryCatch(
    biogeme_bridge()$evaluate_expression_c_biogeme(
      compiled_model = compiled,
      formula_name = "rbiogeme_expression_c",
      beta_values = reticulate::r_to_py(as.list(beta)),
      aggregation = aggregation,
      number_of_draws = as.integer(number_of_draws),
      numerically_safe = numerically_safe,
      use_jit = use_jit
    ),
    error = function(error) {
      biogeme_rethrow(
        error,
        class = "biogeme_evaluation_error",
        operation = "native Biogeme expression calculation",
        suggestion = "check the expression, database, parameter names, and draw count"
      )
    }
  )
  value <- reticulate::py_to_r(raw)$value
  if (isTRUE(aggregation)) return(as.numeric(value)[[1L]])
  as.numeric(unlist(value, use.names = FALSE))
}

#' @export
print.biogeme_simulation <- function(x, ...) {
  cat("Biogeme simulation with ", nrow(x$values), " rows and ", ncol(x$values), " expressions\n", sep = "")
  print(utils::head(x$values))
  invisible(x)
}

#' @export
as.data.frame.biogeme_simulation <- function(x, ...) x$values

#' Validate a model using native Biogeme cross-validation
#' @param model A `biogeme_model`.
#' @param fit A `biogeme_fit` from the same model.
#' @param folds Number of validation folds.
#' @param groups Optional grouping column.
#' @param seed Optional seed controlling native random fold assignment.
#' @param control Optional [biogeme_control()] object.
#' @return A native-bridge validation result.
#' @export
validate <- function(model, fit, folds = 5L, groups = NULL, seed = NULL, control = NULL) {
  if (!inherits(model, "biogeme_model") ||
      (!inherits(fit, "biogeme_fit") && !inherits(fit, "biogeme_bayesian_fit"))) {
    stop("model must be a Biogeme model and fit must be an estimation result.", call. = FALSE)
  }
  if (!is.numeric(folds) || length(folds) != 1L || is.na(folds) || folds < 2 || folds != floor(folds)) {
    stop("folds must be one integer greater than one.", call. = FALSE)
  }
  if (!is.null(seed) && (!is.numeric(seed) || length(seed) != 1L ||
      is.na(seed) || !is.finite(seed) || seed != floor(seed))) {
    stop("seed must be NULL or one finite integer.", call. = FALSE)
  }
  if (!is.null(groups)) {
    groups <- validate_name(groups, "groups")
  }
  controls <- control %||% model$control %||% list()
  compiled <- biogeme_compile_model(model)
  raw <- tryCatch(
    biogeme_bridge()$validate_biogeme(
      compiled_model = compiled,
      beta_values = reticulate::r_to_py(as.list(stats::coef(fit))),
      slices = as.integer(folds),
      groups = groups,
      seed = if (is.null(seed)) NULL else as.integer(seed),
      controls = if (length(controls) == 0L) NULL else reticulate::r_to_py(controls)
    ),
    error = function(error) {
      biogeme_rethrow(
        error,
        class = "biogeme_estimation_error",
        operation = "Biogeme validation",
        suggestion = "check the fold count, grouping column, and model fit"
      )
    }
  )
  reticulate::py_to_r(raw)
}

`%||%` <- function(x, y) if (is.null(x)) y else x

Try the rbiogeme package in your browser

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

rbiogeme documentation built on Sept. 29, 2026, 5:09 p.m.