R/logit_model.R

Defines functions biogeme_native_parameter_names biogeme_model_parameters print.biogeme_logit_model cross_nested_logit_model cross_nested_log_probability cross_nested_probability cross_nested_sparsity_report cross_nested_nests cross_nested_nest cross_nested_logit_correlation nested_logit_correlation nested_logit_model nested_endogenous_sampling_log_probability nested_probability nested_log_probability nested_nests nested_nest logit_model logit_log_probability logit_probability collect_model_betas validate_expression_variables collect_biogeme_betas validate_availability normalize_alternative_codes validate_choice_column validate_alternative_names

Documented in biogeme_model_parameters biogeme_native_parameter_names cross_nested_logit_correlation cross_nested_logit_model cross_nested_log_probability cross_nested_nest cross_nested_nests cross_nested_probability cross_nested_sparsity_report logit_log_probability logit_model logit_probability nested_endogenous_sampling_log_probability nested_logit_correlation nested_logit_model nested_log_probability nested_nest nested_nests nested_probability

validate_alternative_names <- function(utilities) {
  if (!is.list(utilities) || length(utilities) == 0L) {
    stop("utilities must be a non-empty named list.", call. = FALSE)
  }
  utility_names <- names(utilities)
  if (is.null(utility_names) || anyNA(utility_names) || any(!nzchar(utility_names))) {
    stop("utilities must have non-empty names.", call. = FALSE)
  }
  if (anyDuplicated(utility_names)) {
    stop("Utility alternative names must be unique.", call. = FALSE)
  }
  utility_names
}

validate_choice_column <- function(database, choice) {
  validate_biogeme_database(database)
  choice <- validate_name(choice, "choice")
  if (!biogeme_database_has_column(database, choice)) {
    stop(
      "Choice column '", choice, "' is not present in the database.",
      call. = FALSE
    )
  }
  choice_values <- database$data[[choice]]
  if (anyNA(choice_values) || any(!is.finite(choice_values))) {
    stop("The choice column must not contain missing or non-finite values.", call. = FALSE)
  }
  if (any(choice_values != floor(choice_values))) {
    stop("The choice column must contain integer-valued alternatives.", call. = FALSE)
  }
  choice
}

normalize_alternative_codes <- function(
    utility_names,
    choice_values,
    alternative_codes,
    validate_observed = TRUE
) {
  if (is.null(alternative_codes)) {
    numeric_names <- suppressWarnings(as.integer(utility_names))
    if (anyNA(numeric_names) || anyDuplicated(numeric_names)) {
      stop(
        "Provide alternative_codes when utility names are not unique integer codes.",
        call. = FALSE
      )
    }
    codes <- numeric_names
    names(codes) <- utility_names
  } else {
    if (!is.numeric(alternative_codes) || is.null(names(alternative_codes))) {
      stop("alternative_codes must be a named numeric vector.", call. = FALSE)
    }
    if (length(alternative_codes) != length(utility_names) ||
        anyNA(names(alternative_codes)) ||
        any(!nzchar(names(alternative_codes))) ||
        anyDuplicated(names(alternative_codes)) ||
        !setequal(names(alternative_codes), utility_names)) {
      stop(
        "alternative_codes must contain exactly one code for each utility name.",
        call. = FALSE
      )
    }
    if (anyNA(alternative_codes) ||
        any(!is.finite(alternative_codes)) ||
        any(alternative_codes != floor(alternative_codes)) ||
        anyDuplicated(alternative_codes)) {
      stop("Alternative codes must be unique finite integers.", call. = FALSE)
    }
    codes <- as.integer(alternative_codes[utility_names])
    names(codes) <- utility_names
  }

  if (isTRUE(validate_observed)) {
    observed <- unique(as.integer(choice_values))
    if (any(!observed %in% unname(codes))) {
      missing_codes <- paste(setdiff(observed, unname(codes)), collapse = ", ")
      stop(
        "The choice column contains alternative code(s) not present in utilities: ",
        missing_codes,
        call. = FALSE
      )
    }
  }
  codes
}

validate_availability <- function(availability, utility_names) {
  if (is.null(availability)) {
    return(NULL)
  }
  if (!is.list(availability) || is.null(names(availability)) ||
      length(availability) != length(utility_names) ||
      anyDuplicated(names(availability)) ||
      !setequal(names(availability), utility_names)) {
    stop(
      "availability must contain exactly one expression for each utility name.",
      call. = FALSE
    )
  }
  availability <- availability[utility_names]
  result <- lapply(availability, as_biogeme_expression)
  for (name in utility_names) {
    expression <- result[[name]]
    if (identical(expression$kind, "numeric") &&
        !expression$value %in% c(0, 1)) {
      stop(
        "Availability for '", name, "' must be 0 or 1 when it is numeric.",
        call. = FALSE
      )
    }
  }
  result
}

collect_biogeme_betas <- function(expression, collected = list()) {
  if (!is_biogeme_expression(expression)) {
    stop("Internal error: expected a Biogeme expression.", call. = FALSE)
  }
  if (identical(expression$kind, "beta")) {
    name <- expression$name
    if (!is.null(collected[[name]])) {
      existing <- collected[[name]]
      fields <- c("name", "start", "lower", "upper", "fixed")
      if (!identical(existing[fields], expression[fields])) {
        stop(
          "Parameter '", name,
          "' is specified inconsistently in the model.",
          call. = FALSE
        )
      }
    } else {
      collected[[name]] <- expression
    }
    return(collected)
  }
  if (identical(expression$kind, "segmented_beta")) {
    for (definition in biogeme_segmented_beta_definitions(expression)) {
      collected <- collect_biogeme_betas(definition, collected)
    }
    return(collected)
  }
  for (child in biogeme_expression_children(expression)) {
    collected <- collect_biogeme_betas(child, collected)
  }
  collected
}

validate_expression_variables <- function(expressions, database, argument = "expression") {
  known <- biogeme_database_columns(database)
  if (length(expressions) == 0L) return(invisible(NULL))
  for (index in seq_along(expressions)) {
    expression <- expressions[[index]]
    missing <- setdiff(collect_biogeme_variables(expression), known)
    if (length(missing) > 0L) {
      stop(
        argument, " refers to unknown database variable(s): ",
        paste(missing, collapse = ", "),
        call. = FALSE
      )
    }
  }
  invisible(NULL)
}

collect_model_betas <- function(utilities, availability, weight) {
  collected <- list()
  for (expression in utilities) {
    collected <- collect_biogeme_betas(expression, collected)
  }
  if (!is.null(availability)) {
    for (expression in availability) {
      collected <- collect_biogeme_betas(expression, collected)
    }
  }
  if (!is.null(weight)) {
    collected <- collect_biogeme_betas(weight, collected)
  }
  collected
}

#' Construct a native logit probability expression
#'
#' @param utilities Named list of utility expressions, one per alternative.
#' @param availability Optional named list of availability expressions.
#' @param alternative Integer-valued alternative whose probability is returned,
#'   or a Biogeme expression selecting the alternative for each observation.
#' @param alternative_codes Optional named integer vector when utility names
#'   are not themselves integer alternative codes.
#' @return A probability expression compiled to native `biogeme.models.logit`.
#' @export
logit_probability <- function(
    utilities,
    availability = NULL,
    alternative,
    alternative_codes = NULL
) {
  utility_names <- validate_alternative_names(utilities)
  utilities <- lapply(utilities, as_biogeme_expression)
  names(utilities) <- utility_names
  availability <- validate_availability(availability, utility_names)

  alternative_is_expression <- is_biogeme_expression(alternative)
  if (alternative_is_expression) {
    alternative <- as_biogeme_expression(alternative)
  } else if (!is.numeric(alternative) || length(alternative) != 1L ||
             is.na(alternative) || !is.finite(alternative) ||
             alternative != floor(alternative)) {
    stop("alternative must be one finite integer code or a Biogeme expression.", call. = FALSE)
  }
  if (is.null(alternative_codes)) {
    codes <- suppressWarnings(as.integer(utility_names))
    if (anyNA(codes) || anyDuplicated(codes)) {
      stop(
        "Provide alternative_codes when utility names are not integer codes.",
        call. = FALSE
      )
    }
    names(codes) <- utility_names
  } else {
    if (!is.numeric(alternative_codes) || is.null(names(alternative_codes)) ||
        length(alternative_codes) != length(utility_names) ||
        anyNA(names(alternative_codes)) || any(!nzchar(names(alternative_codes))) ||
        anyDuplicated(names(alternative_codes)) ||
        !setequal(names(alternative_codes), utility_names) ||
        anyNA(alternative_codes) || any(!is.finite(alternative_codes)) ||
        any(alternative_codes != floor(alternative_codes)) ||
        anyDuplicated(alternative_codes)) {
      stop("alternative_codes must contain unique finite integer codes.", call. = FALSE)
    }
    codes <- as.integer(alternative_codes[utility_names])
    names(codes) <- utility_names
  }
  if (!alternative_is_expression) alternative <- as.integer(alternative)
  if (!alternative_is_expression && !alternative %in% unname(codes)) {
    stop("alternative must be one of the utility alternative codes.", call. = FALSE)
  }

  new_biogeme_expression(
    "logit_probability",
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = codes
  )
}

#' Construct a native logit log-probability expression
#'
#' This is the log-probability counterpart of [logit_probability()]. It is
#' compiled directly to Biogeme's public `models.loglogit` constructor, which
#' preserves the native numerically stable likelihood used by examples such
#' as Swissmetro b20.
#'
#' @param utilities Named list of utility expressions keyed by alternatives.
#' @param availability Optional named list of availability expressions.
#' @param alternative Integer code or expression for the observed alternative.
#' @param alternative_codes Optional named integer codes for the utilities.
#' @return A Biogeme logit log-probability expression.
#' @export
logit_log_probability <- function(
    utilities,
    availability = NULL,
    alternative,
    alternative_codes = NULL
) {
  expression <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  expression$kind <- "logit_log_probability"
  expression
}

#' @rdname logit_probability
#' @export
logit <- logit_probability

#' Define a cross-sectional multinomial logit model
#'
#' @param database A `biogeme_database` object.
#' @param choice Name of the integer-valued choice column.
#' @param utilities Named list of utility expressions, one per alternative.
#' @param availability Optional named list of availability expressions.
#' @param alternative_codes Optional named integer vector mapping utility names
#'   to values in the choice column. If omitted, utility names must themselves
#'   be integer codes.
#' @param weight Optional observation-weight expression.
#' @return An object of class `biogeme_logit_model`.
#' @details
#' The utility list is named by alternative. If the names are not integer
#' codes, pass `alternative_codes` as a named integer vector. Availability and
#' weights may be constants or symbolic expressions.
#' @examples
#' database <- biogeme_database(
#'   "demo",
#'   data.frame(choice = c(1, 2), x = c(1, 2))
#' )
#' b <- biogeme_beta("b")
#' model <- logit_model(
#'   database,
#'   choice = "choice",
#'   utilities = list(`1` = 0, `2` = b * variable("x"))
#' )
#' model
#' @export
logit_model <- function(
    database,
    choice,
    utilities,
    availability = NULL,
    alternative_codes = NULL,
    weight = NULL
) {
  validate_biogeme_database(database)
  choice <- validate_choice_column(database, choice)
  utility_names <- validate_alternative_names(utilities)
  utilities <- lapply(utilities, as_biogeme_expression)
  names(utilities) <- utility_names
  availability <- validate_availability(availability, utility_names)
  if (!is.null(weight)) {
    weight <- as_biogeme_expression(weight)
    if (identical(weight$kind, "numeric") && weight$value < 0) {
      stop("A numeric observation weight must not be negative.", call. = FALSE)
    }
  }

  codes <- normalize_alternative_codes(
    utility_names,
    database$data[[choice]],
    alternative_codes,
    validate_observed = length(database$filters) == 0L
  )
  parameters <- collect_model_betas(utilities, availability, weight)
  validate_expression_variables(
    c(unname(utilities), if (is.null(availability)) list() else unname(availability),
      if (is.null(weight)) list() else list(weight)),
    database,
    argument = "model expression"
  )

  structure(
    list(
      kind = "logit",
      database = database,
      choice = choice,
      choice_expression = variable(choice),
      utilities = utilities,
      availability = availability,
      alternative_codes = codes,
      weight = weight,
      parameters = parameters
    ),
  class = c("biogeme_logit_model", "biogeme_model")
  )
}

#' Define one non-trivial nested-logit nest
#'
#' @param nest_parameter Native nest parameter expression or numeric value.
#' @param alternatives Integer alternative codes contained in the nest.
#' @param name Optional native nest name.
#' @return A `biogeme_nested_nest` specification.
#' @export
nested_nest <- function(nest_parameter, alternatives, name = NULL) {
  nest_parameter <- as_biogeme_expression(nest_parameter)
  if (!is.numeric(alternatives) || length(alternatives) == 0L ||
      anyNA(alternatives) || any(!is.finite(alternatives)) ||
      any(alternatives != floor(alternatives)) || anyDuplicated(alternatives)) {
    stop("alternatives must contain unique finite integer codes.", call. = FALSE)
  }
  alternatives <- as.integer(alternatives)
  if (!is.null(name)) name <- validate_name(name, "name")
  structure(
    list(
      nest_parameter = nest_parameter,
      alternatives = alternatives,
      name = name
    ),
    class = "biogeme_nested_nest"
  )
}

#' Define the complete nested-logit nest structure
#'
#' Alternatives not listed in a non-trivial nest are treated by native
#' Biogeme as trivial nests containing one alternative.
#'
#' @param choice_set All integer alternative codes.
#' @param nests A non-empty list of [nested_nest()] specifications.
#' @return A `biogeme_nested_nests` specification.
#' @export
nested_nests <- function(choice_set, nests) {
  if (!is.numeric(choice_set) || length(choice_set) == 0L ||
      anyNA(choice_set) || any(!is.finite(choice_set)) ||
      any(choice_set != floor(choice_set)) || anyDuplicated(choice_set)) {
    stop("choice_set must contain unique finite integer codes.", call. = FALSE)
  }
  choice_set <- as.integer(choice_set)
  if (!is.list(nests) || length(nests) == 0L ||
      any(!vapply(nests, inherits, logical(1), what = "biogeme_nested_nest"))) {
    stop("nests must be a non-empty list of nested_nest() specifications.", call. = FALSE)
  }
  listed <- unlist(lapply(nests, function(nest) nest$alternatives), use.names = FALSE)
  if (any(!listed %in% choice_set)) {
    stop("Every nested alternative must belong to choice_set.", call. = FALSE)
  }
  if (anyDuplicated(listed)) {
    stop("An alternative cannot belong to more than one non-trivial nest.", call. = FALSE)
  }
  names <- vapply(nests, function(nest) {
    if (is.null(nest$name)) "" else nest$name
  }, character(1))
  if (anyDuplicated(names[nzchar(names)])) {
    stop("Nested-logit nest names must be unique.", call. = FALSE)
  }
  structure(
    list(choice_set = choice_set, nests = nests),
    class = "biogeme_nested_nests"
  )
}

#' Construct a native nested-logit log-probability expression
#'
#' @param utilities Named list of utility expressions.
#' @param availability Optional named list of availability expressions.
#' @param nests A [nested_nests()] specification.
#' @param alternative Integer-valued alternative or a Biogeme choice expression.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @param scale_parameter Optional native scale expression for bottom normalization.
#' @return A Biogeme expression compiled to native `biogeme.models.lognested`.
#' @export
nested_log_probability <- function(
    utilities,
    availability,
    nests,
    alternative,
    alternative_codes = NULL,
    scale_parameter = NULL
) {
  if (!inherits(nests, "biogeme_nested_nests")) {
    stop("nests must be a nested_nests() specification.", call. = FALSE)
  }
  validated <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  if (!setequal(unname(validated$alternative_codes), nests$choice_set)) {
    stop("nests choice_set must match the utility alternative codes.", call. = FALSE)
  }
  if (!is.null(scale_parameter)) scale_parameter <- as_biogeme_expression(scale_parameter)
  kind <- if (is.null(scale_parameter)) "nested_logit" else "nested_logit_mev_mu"
  new_biogeme_expression(
    kind,
    utilities = validated$utilities,
    availability = validated$availability,
    nests = nests,
    alternative = validated$alternative,
    alternative_codes = validated$alternative_codes,
    scale_parameter = scale_parameter
  )
}

#' Construct a native nested-logit probability expression
#'
#' @param utilities Named list of utility expressions.
#' @param availability Optional named list of availability expressions.
#' @param nests A [nested_nests()] specification.
#' @param alternative Integer-valued alternative or a Biogeme choice expression.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @return A Biogeme expression compiled to native `biogeme.models.nested`.
#' @export
nested_probability <- function(
    utilities,
    availability,
    nests,
    alternative,
    alternative_codes = NULL
) {
  if (!inherits(nests, "biogeme_nested_nests")) {
    stop("nests must be a nested_nests() specification.", call. = FALSE)
  }
  validated <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  if (!setequal(unname(validated$alternative_codes), nests$choice_set)) {
    stop("nests choice_set must match the utility alternative codes.", call. = FALSE)
  }
  new_biogeme_expression(
    "nested_probability",
    utilities = validated$utilities,
    availability = validated$availability,
    nests = nests,
    alternative = validated$alternative,
    alternative_codes = validated$alternative_codes
  )
}

#' Construct a native nested-logit log probability with endogenous-sampling correction
#'
#' @param utilities Named list of utility expressions.
#' @param availability Optional named list of availability expressions.
#' @param nests A [nested_nests()] specification.
#' @param correction Named list of alternative-specific log correction terms.
#' @param alternative Integer-valued alternative or a Biogeme choice expression.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @return A Biogeme expression compiled through native `get_mev_for_nested`
#'   and `logmev_endogenous_sampling`.
#' @export
nested_endogenous_sampling_log_probability <- function(
    utilities,
    availability,
    nests,
    correction,
    alternative,
    alternative_codes = NULL
) {
  if (!inherits(nests, "biogeme_nested_nests")) {
    stop("nests must be a nested_nests() specification.", call. = FALSE)
  }
  validated <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  utility_names <- names(validated$utilities)
  if (!is.list(correction) || is.null(names(correction)) ||
      length(correction) != length(utility_names) ||
      anyNA(names(correction)) || any(!nzchar(names(correction))) ||
      anyDuplicated(names(correction)) || !setequal(names(correction), utility_names)) {
    stop("correction must contain exactly one expression for each utility.", call. = FALSE)
  }
  correction <- lapply(correction[utility_names], as_biogeme_expression)
  names(correction) <- utility_names
  if (!setequal(unname(validated$alternative_codes), nests$choice_set)) {
    stop("nests choice_set must match the utility alternative codes.", call. = FALSE)
  }
  new_biogeme_expression(
    "nested_endogenous_sampling",
    utilities = validated$utilities,
    availability = validated$availability,
    nests = nests,
    correction = correction,
    alternative = validated$alternative,
    alternative_codes = validated$alternative_codes
  )
}

#' Define a cross-sectional nested-logit model
#'
#' @param database A `biogeme_database` object.
#' @param choice Name of the integer-valued choice column.
#' @param utilities Named list of utility expressions.
#' @param nests A [nested_nests()] specification.
#' @param availability Optional named list of availability expressions.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @param scale_parameter Optional native scale expression for bottom normalization.
#' @param control Optional [biogeme_control()] object.
#' @return An object of class `biogeme_nested_logit_model`.
#' @export
nested_logit_model <- function(
    database,
    choice,
    utilities,
    nests,
    availability = NULL,
    alternative_codes = NULL,
    scale_parameter = NULL,
    control = NULL
) {
  validate_biogeme_database(database)
  choice <- validate_choice_column(database, choice)
  log_probability <- nested_log_probability(
    utilities = utilities,
    availability = availability,
    nests = nests,
    alternative = variable(choice),
    alternative_codes = alternative_codes,
    scale_parameter = scale_parameter
  )
  model <- biogeme_model(
    database = database,
    formula = log_probability,
    control = control
  )
  model$kind <- "nested_logit"
  model$choice <- choice
  model$choice_expression <- variable(choice)
  model$utilities <- log_probability$utilities
  model$availability <- log_probability$availability
  model$alternative_codes <- log_probability$alternative_codes
  model$nests <- nests
  model$scale_parameter <- scale_parameter
  class(model) <- c("biogeme_nested_logit_model", "biogeme_model")
  model
}

#' Calculate the native nested-logit error-term correlation matrix
#'
#' @param model A `biogeme_nested_logit_model`.
#' @param beta_values Optional named estimated or starting parameter values.
#' @param alternatives_names Optional named character vector for row/column labels.
#' @param mu Overall scale parameter passed to native Biogeme.
#' @return A numeric correlation matrix.
#' @export
nested_logit_correlation <- function(
    model,
    beta_values = NULL,
    alternatives_names = NULL,
    mu = 1
) {
  if (!inherits(model, "biogeme_nested_logit_model")) {
    stop("model must be a nested_logit_model() object.", call. = FALSE)
  }
  if (is.null(beta_values)) {
    beta_values <- vapply(model$parameters, function(parameter) parameter$start, numeric(1))
    names(beta_values) <- names(model$parameters)
  }
  if (!is.numeric(beta_values) || is.null(names(beta_values)) ||
      anyNA(names(beta_values)) || any(!nzchar(names(beta_values))) ||
      anyNA(beta_values) || any(!is.finite(beta_values))) {
    stop("beta_values must be a named finite numeric vector.", call. = FALSE)
  }
  if (!is.numeric(mu) || length(mu) != 1L || is.na(mu) || !is.finite(mu) || mu <= 0) {
    stop("mu must be one positive finite numeric value.", call. = FALSE)
  }
  if (is.null(alternatives_names)) {
    alternatives_names <- stats::setNames(
      as.character(model$nests$choice_set),
      as.character(model$nests$choice_set)
    )
  }
  if (!is.character(alternatives_names) || is.null(names(alternatives_names)) ||
      anyNA(names(alternatives_names)) || any(!nzchar(names(alternatives_names))) ||
      anyNA(alternatives_names) || any(!nzchar(alternatives_names))) {
    stop("alternatives_names must be a named non-empty character vector.", call. = FALSE)
  }
  compiled <- biogeme_compile_model(model)
  bridge <- biogeme_bridge()
  converted <- reticulate::py_to_r(bridge$nested_logit_correlation(
    nests = compiled$nests,
    parameters = reticulate::r_to_py(as.list(beta_values)),
    alternatives_names = reticulate::r_to_py(as.list(alternatives_names)),
    mu = as.numeric(mu)
  ))
  values <- converted$data
  if (!is.matrix(values)) values <- do.call(rbind, values)
  values <- unname(values)
  dimnames(values) <- list(as.character(converted$index), as.character(converted$columns))
  values
}

#' Calculate the native cross-nested-logit error-term correlation matrix
#'
#' @param model A `biogeme_cross_nested_logit_model`.
#' @param beta_values Optional named estimated or starting parameter values.
#' @param alternatives_names Optional named character vector for row/column labels.
#' @return A numeric correlation matrix.
#' @export
cross_nested_logit_correlation <- function(
    model,
    beta_values = NULL,
    alternatives_names = NULL
) {
  if (!inherits(model, "biogeme_cross_nested_logit_model")) {
    stop("model must be a cross_nested_logit_model() object.", call. = FALSE)
  }
  if (is.null(beta_values)) {
    beta_values <- vapply(model$parameters, function(parameter) parameter$start, numeric(1))
    names(beta_values) <- names(model$parameters)
  }
  if (!is.numeric(beta_values) || is.null(names(beta_values)) ||
      anyNA(names(beta_values)) || any(!nzchar(names(beta_values))) ||
      anyNA(beta_values) || any(!is.finite(beta_values))) {
    stop("beta_values must be a named finite numeric vector.", call. = FALSE)
  }
  if (is.null(alternatives_names)) {
    alternatives_names <- stats::setNames(
      as.character(model$nests$choice_set),
      as.character(model$nests$choice_set)
    )
  }
  if (!is.character(alternatives_names) || is.null(names(alternatives_names)) ||
      anyNA(names(alternatives_names)) || any(!nzchar(names(alternatives_names))) ||
      anyNA(alternatives_names) || any(!nzchar(alternatives_names))) {
    stop("alternatives_names must be a named non-empty character vector.", call. = FALSE)
  }
  compiled <- biogeme_compile_model(model)
  bridge <- biogeme_bridge()
  converted <- reticulate::py_to_r(bridge$cross_nested_logit_correlation(
    nests = compiled$nests,
    parameters = reticulate::r_to_py(as.list(beta_values)),
    alternatives_names = reticulate::r_to_py(as.list(alternatives_names))
  ))
  values <- converted$data
  if (!is.matrix(values)) values <- do.call(rbind, values)
  values <- unname(values)
  dimnames(values) <- list(as.character(converted$index), as.character(converted$columns))
  values
}

#' Define one cross-nested-logit nest
#'
#' @param nest_parameter Native nest parameter expression or numeric value.
#' @param allocation Named mapping from alternative codes to allocation
#'   expressions in this nest.
#' @param name Optional native nest name.
#' @return A `biogeme_cross_nested_nest` specification.
#' @export
cross_nested_nest <- function(nest_parameter, allocation, name = NULL) {
  nest_parameter <- as_biogeme_expression(nest_parameter)
  if (is.atomic(allocation) && !is.null(names(allocation))) {
    allocation <- as.list(allocation)
  }
  if (!is.list(allocation) || is.null(names(allocation)) || length(allocation) == 0L ||
      anyNA(names(allocation)) || any(!nzchar(names(allocation)))) {
    stop("allocation must be a non-empty named mapping.", call. = FALSE)
  }
  keys <- suppressWarnings(as.integer(names(allocation)))
  if (anyNA(keys) || anyDuplicated(keys)) {
    stop("allocation keys must be unique integer alternative codes.", call. = FALSE)
  }
  allocation <- lapply(allocation, as_biogeme_expression)
  for (expression in allocation) {
    if (identical(expression$kind, "numeric") &&
        (expression$value < 0 || expression$value > 1)) {
      stop("Numeric allocation values must lie in [0, 1].", call. = FALSE)
    }
  }
  names(allocation) <- as.character(keys)
  if (!is.null(name)) name <- validate_name(name, "name")
  structure(
    list(nest_parameter = nest_parameter, allocation = allocation, name = name),
    class = "biogeme_cross_nested_nest"
  )
}

#' Define the complete cross-nested-logit nest structure
#'
#' @param choice_set All integer alternative codes.
#' @param nests A non-empty list of [cross_nested_nest()] specifications.
#' @param sparse If `TRUE`, allow structurally zero allocation entries to be
#'   omitted. Supplied symbolic allocations are always retained.
#' @return A `biogeme_cross_nested_nests` specification.
#' @export
cross_nested_nests <- function(choice_set, nests, sparse = FALSE) {
  if (!is.numeric(choice_set) || length(choice_set) == 0L ||
      anyNA(choice_set) || any(!is.finite(choice_set)) ||
      any(choice_set != floor(choice_set)) || anyDuplicated(choice_set)) {
    stop("choice_set must contain unique finite integer codes.", call. = FALSE)
  }
  choice_set <- as.integer(choice_set)
  if (!is.logical(sparse) || length(sparse) != 1L || is.na(sparse)) {
    stop("sparse must be one non-missing logical value.", call. = FALSE)
  }
  if (!is.list(nests) || length(nests) == 0L ||
      any(!vapply(nests, inherits, logical(1), what = "biogeme_cross_nested_nest"))) {
    stop("nests must be a non-empty list of cross_nested_nest() specifications.", call. = FALSE)
  }
  for (nest in nests) {
    allocation_keys <- as.integer(names(nest$allocation))
    if (isTRUE(sparse)) {
      if (any(!allocation_keys %in% choice_set)) {
        stop("Sparse allocations must contain only choice_set alternatives.", call. = FALSE)
      }
    } else if (!setequal(allocation_keys, choice_set)) {
      stop("Every cross-nested allocation must cover choice_set exactly.", call. = FALSE)
    }
  }
  names <- vapply(nests, function(nest) {
    if (is.null(nest$name)) "" else nest$name
  }, character(1))
  if (anyDuplicated(names[nzchar(names)])) {
    stop("Cross-nested nest names must be unique.", call. = FALSE)
  }
  structure(
    list(choice_set = choice_set, nests = nests, sparse = isTRUE(sparse)),
    class = "biogeme_cross_nested_nests"
  )
}

#' Summarize the structural sparsity of a CNL nest specification
#'
#' @param nests A [cross_nested_nests()] specification.
#' @return A data frame with one row per nest and membership counts.
#' @export
cross_nested_sparsity_report <- function(nests) {
  if (!inherits(nests, "biogeme_cross_nested_nests")) {
    stop("nests must be a cross_nested_nests() specification.", call. = FALSE)
  }
  possible <- length(nests$choice_set)
  nest_names <- vapply(seq_along(nests$nests), function(index) {
    name <- nests$nests[[index]]$name
    if (is.null(name)) paste0("nest", index) else name
  }, character(1))
  stored <- vapply(nests$nests, function(nest) length(nest$allocation), integer(1))
  zero <- vapply(nests$nests, function(nest) {
    sum(vapply(nest$allocation, function(expression) {
      identical(expression$kind, "numeric") && identical(expression$value, 0)
    }, logical(1)))
  }, integer(1))
  active <- stored - zero
  data.frame(
    nest = nest_names,
    possible_memberships = rep.int(possible, length(nests$nests)),
    stored_memberships = stored,
    active_memberships = active,
    omitted_memberships = possible - stored,
    zero_allocations = zero,
    density = active / possible,
    sparse = isTRUE(nests$sparse),
    stringsAsFactors = FALSE
  )
}

#' Construct a native cross-nested-logit probability expression
#'
#' @param utilities Named list of utility expressions.
#' @param availability Optional named list of availability expressions.
#' @param nests A [cross_nested_nests()] specification.
#' @param alternative Integer-valued alternative or a Biogeme choice expression.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @return A Biogeme expression compiled to native `biogeme.models.cnl`.
#' @export
cross_nested_probability <- function(
    utilities,
    availability,
    nests,
    alternative,
    alternative_codes = NULL
) {
  if (!inherits(nests, "biogeme_cross_nested_nests")) {
    stop("nests must be a cross_nested_nests() specification.", call. = FALSE)
  }
  validated <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  if (!setequal(unname(validated$alternative_codes), nests$choice_set)) {
    stop("nests choice_set must match the utility alternative codes.", call. = FALSE)
  }
  new_biogeme_expression(
    "cross_nested_probability",
    utilities = validated$utilities,
    availability = validated$availability,
    nests = nests,
    alternative = validated$alternative,
    alternative_codes = validated$alternative_codes
  )
}

#' Construct a native cross-nested-logit log-probability expression
#'
#' @param utilities Named list of utility expressions.
#' @param availability Optional named list of availability expressions.
#' @param nests A [cross_nested_nests()] specification.
#' @param alternative Integer-valued alternative or a Biogeme choice expression.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @param scale_parameter Optional native CNL scale expression.
#' @return A Biogeme expression compiled to native `biogeme.models.logcnl`.
#' @export
cross_nested_log_probability <- function(
    utilities,
    availability,
    nests,
    alternative,
    alternative_codes = NULL,
    scale_parameter = NULL
) {
  if (!inherits(nests, "biogeme_cross_nested_nests")) {
    stop("nests must be a cross_nested_nests() specification.", call. = FALSE)
  }
  validated <- logit_probability(
    utilities = utilities,
    availability = availability,
    alternative = alternative,
    alternative_codes = alternative_codes
  )
  if (!setequal(unname(validated$alternative_codes), nests$choice_set)) {
    stop("nests choice_set must match the utility alternative codes.", call. = FALSE)
  }
  if (!is.null(scale_parameter)) scale_parameter <- as_biogeme_expression(scale_parameter)
  kind <- if (is.null(scale_parameter)) "cross_nested_logit" else "cross_nested_logit_mu"
  new_biogeme_expression(
    kind,
    utilities = validated$utilities,
    availability = validated$availability,
    nests = nests,
    alternative = validated$alternative,
    alternative_codes = validated$alternative_codes,
    scale_parameter = scale_parameter
  )
}

#' Define a cross-sectional cross-nested-logit model
#'
#' @param database A `biogeme_database` object.
#' @param choice Name of the integer-valued choice column.
#' @param utilities Named list of utility expressions.
#' @param nests A [cross_nested_nests()] specification.
#' @param availability Optional named list of availability expressions.
#' @param alternative_codes Optional named integer vector mapping utility names.
#' @param scale_parameter Optional native CNL scale expression.
#' @param control Optional [biogeme_control()] object.
#' @return An object of class `biogeme_cross_nested_logit_model`.
#' @export
cross_nested_logit_model <- function(
    database,
    choice,
    utilities,
    nests,
    availability = NULL,
    alternative_codes = NULL,
    scale_parameter = NULL,
    control = NULL
) {
  validate_biogeme_database(database)
  choice <- validate_choice_column(database, choice)
  log_probability <- cross_nested_log_probability(
    utilities = utilities,
    availability = availability,
    nests = nests,
    alternative = variable(choice),
    alternative_codes = alternative_codes,
    scale_parameter = scale_parameter
  )
  model <- biogeme_model(
    database = database,
    formula = log_probability,
    control = control
  )
  model$kind <- "cross_nested_logit"
  model$choice <- choice
  model$choice_expression <- variable(choice)
  model$utilities <- log_probability$utilities
  model$availability <- log_probability$availability
  model$alternative_codes <- log_probability$alternative_codes
  model$nests <- nests
  model$scale_parameter <- scale_parameter
  class(model) <- c("biogeme_cross_nested_logit_model", "biogeme_model")
  model
}

#' @export
print.biogeme_logit_model <- function(x, ...) {
  cat(
    "Biogeme cross-sectional logit model with ",
    length(x$utilities), " alternatives and ",
    length(x$parameters), " free/fixed parameter definition(s)\n",
    sep = ""
  )
  invisible(x)
}

#' Return the parameter definitions in a model
#'
#' @param model A Biogeme model.
#' @return A data frame of parameter definitions.
#' @export
biogeme_model_parameters <- function(model) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  if (length(model$parameters) == 0L) {
    return(data.frame(
      name = character(),
      start = numeric(),
      lower = numeric(),
      upper = numeric(),
      fixed = logical(),
      stringsAsFactors = FALSE
    ))
  }
  do.call(
    rbind,
    lapply(model$parameters, function(parameter) {
      data.frame(
        name = parameter$name,
        start = parameter$start,
        lower = if (is.null(parameter$lower)) NA_real_ else parameter$lower,
        upper = if (is.null(parameter$upper)) NA_real_ else parameter$upper,
        fixed = parameter$fixed,
        stringsAsFactors = FALSE
      )
    })
  )
}

#' Return parameter names collected from the compiled native expression
#'
#' Unlike [biogeme_model_parameters()], this operation also sees parameters
#' generated by native helpers such as segmentation and parameters in every
#' catalog branch.  When `configuration_id` is supplied, native Biogeme first
#' resolves that configuration and the returned names describe the flattened
#' selected expression.
#'
#' @param model A Biogeme model.
#' @param configuration_id Optional exact native catalog configuration ID.
#' @param model_name Temporary native model name used while compiling.
#' @param controls Named native Biogeme controls.
#' @return A list with `all`, `free`, and `fixed` character vectors.
#' @export
biogeme_native_parameter_names <- function(
    model,
    configuration_id = NULL,
    model_name = "rbiogeme_parameters",
    controls = list()
) {
  if (!inherits(model, "biogeme_model")) {
    stop("model must be a Biogeme model.", call. = FALSE)
  }
  if (!is.null(configuration_id) &&
      (!is.character(configuration_id) || length(configuration_id) != 1L ||
       is.na(configuration_id) || !nzchar(configuration_id))) {
    stop("configuration_id must be NULL or one non-empty character string.", call. = FALSE)
  }
  if (!is.character(model_name) || length(model_name) != 1L ||
      is.na(model_name) || !nzchar(model_name)) {
    stop("model_name must be one non-empty character string.", call. = FALSE)
  }
  controls <- validate_estimation_controls(controls)
  raw <- biogeme_native_parameter_names_model(
    model = model,
    configuration_id = configuration_id,
    model_name = model_name,
    controls = controls
  )
  lapply(raw, as.character)
}

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.