Nothing
.check_columns <- function(data, cols, arg = "columns") {
missing <- setdiff(cols, names(data))
if (length(missing)) {
stop(
"`", arg, "` must exist in `data`: ",
paste(missing, collapse = ", "),
call. = FALSE
)
}
}
.metric_data <- function(data, cols) {
.check_columns(data, cols)
out <- data[stats::complete.cases(data[, cols, drop = FALSE]), cols, drop = FALSE]
numeric_cols <- cols[vapply(out[, cols, drop = FALSE], is.numeric, logical(1))]
if (length(numeric_cols)) {
finite_rows <- Reduce(
`&`,
lapply(out[, numeric_cols, drop = FALSE], is.finite)
)
out <- out[finite_rows, , drop = FALSE]
}
out
}
.observed_category_counts <- function(x) {
x <- x[!is.na(x)]
if (is.factor(x)) {
x <- droplevels(x)
}
table(x)
}
.observed_n_categories <- function(x) {
length(.observed_category_counts(x))
}
.two_levels <- function(x, arg = "category") {
x <- x[!is.na(x)]
if (is.factor(x)) {
x <- as.character(droplevels(x))
}
levs <- unique(x)
if (length(levs) != 2L) {
stop("`", arg, "` must have exactly two non-missing values.", call. = FALSE)
}
levs
}
.check_numeric_features <- function(data, features) {
if (!length(features)) {
stop("`features` must contain at least one column.", call. = FALSE)
}
is_num <- vapply(data[, features, drop = FALSE], is.numeric, logical(1))
if (!all(is_num)) {
stop(
"All `features` must be numeric. Non-numeric feature(s): ",
paste(features[!is_num], collapse = ", "),
call. = FALSE
)
}
}
.check_conf_level <- function(conf_level) {
if (!is.numeric(conf_level) || length(conf_level) != 1L ||
!is.finite(conf_level) || conf_level <= 0 || conf_level >= 1) {
stop("`conf_level` must be a single finite number between 0 and 1.", call. = FALSE)
}
}
.check_positive_count <- function(x, arg) {
if (!is.numeric(x) || length(x) != 1L || !is.finite(x) ||
x < 1 || x != as.integer(x)) {
stop("`", arg, "` must be a positive integer.", call. = FALSE)
}
}
.check_group_cols <- function(group_col, arg = "group_col") {
if (is.null(group_col)) {
return(NULL)
}
if (!is.character(group_col) || !length(group_col) ||
anyNA(group_col) || any(!nzchar(group_col))) {
stop("`", arg, "` must be a character vector of one or more column names.", call. = FALSE)
}
if (anyDuplicated(group_col)) {
stop("`", arg, "` must not contain duplicate column names.", call. = FALSE)
}
group_col
}
.validate_metric_inputs <- function(data, features, category_col, group_col = NULL) {
if (!is.data.frame(data)) {
stop("`data` must be a data frame.", call. = FALSE)
}
if (!is.character(category_col) || length(category_col) != 1L ||
is.na(category_col) || !nzchar(category_col)) {
stop("`category_col` must be a single column name.", call. = FALSE)
}
if (!is.character(features) || !length(features) || anyNA(features) ||
any(!nzchar(features))) {
stop("`features` must be a non-empty character vector of column names.", call. = FALSE)
}
group_col <- .check_group_cols(group_col)
clash <- intersect(features, c(category_col, group_col))
if (length(clash)) {
stop(
"`features` must not overlap with `category_col`/`group_col`: ",
paste(clash, collapse = ", "),
call. = FALSE
)
}
if (!is.null(group_col) && category_col %in% group_col) {
stop("`category_col` must not also appear in `group_col`.", call. = FALSE)
}
invisible(TRUE)
}
.warn_failed_groups <- function(out, value_col, fn) {
if (nrow(out) > 0L && value_col %in% names(out)) {
failed <- sum(is.na(out[[value_col]]))
if (failed) {
warning(
fn, ": ", failed, " of ", nrow(out),
" group(s) could not be estimated and were returned as NA.",
call. = FALSE
)
}
}
out
}
.group_label <- function(data, group_col) {
group_col <- .check_group_cols(group_col)
if (length(group_col) == 1L) {
return(data[[group_col]][1])
}
vals <- vapply(group_col, function(col) {
as.character(data[[col]][1])
}, character(1))
paste(paste0(group_col, "=", vals), collapse = " | ")
}
.match_kde_engine <- function(engine) {
engine <- match.arg(engine, c("ks", "fast_diag", "fast_diagonal"))
if (identical(engine, "fast_diagonal")) {
return("fast_diag")
}
engine
}
.check_ridge_eps <- function(eps, arg = "eps") {
if (!is.numeric(eps) || length(eps) != 1L || !is.finite(eps) || eps <= 0) {
stop("`", arg, "` must be a single positive finite number.", call. = FALSE)
}
}
.empty_group_pillai <- function() {
tibble::tibble(
group = character(),
n_tokens = integer(),
pillai = numeric(),
p_value = numeric()
)
}
.empty_group_jsd <- function() {
tibble::tibble(group = character(), n_tokens = integer(), jsd = numeric())
}
.empty_group_bhatt <- function() {
data.frame(
group = character(),
n_tokens = integer(),
bhatt_dist = numeric(),
bhatt_affinity = numeric(),
stringsAsFactors = FALSE
)
}
.empty_estimate_bhatt_group <- function() {
data.frame(
scope = character(),
group = character(),
n_tokens = integer(),
bhatt_dist = numeric(),
bhatt_affinity = numeric(),
stringsAsFactors = FALSE
)
}
.kde_min_category_tokens <- function(n_features) {
max(2L, n_features + 1L)
}
.check_two_category_sample_size <- function(data, category_col, min_per_category, metric) {
counts <- .observed_category_counts(data[[category_col]])
if (length(counts) != 2L || any(counts < min_per_category)) {
stop(
metric, " requires at least ", min_per_category,
" finite observations in each category after removing missing values.",
call. = FALSE
)
}
}
.split_groups <- function(data, group_col) {
group_col <- .check_group_cols(group_col)
if (length(group_col) == 1L) {
return(split(data, data[[group_col]], drop = TRUE))
}
group_key <- interaction(data[, group_col, drop = FALSE], drop = TRUE, sep = "\r")
groups <- split(data, group_key, drop = TRUE)
names(groups) <- vapply(groups, .group_label, character(1), group_col = group_col)
groups
}
.select_univariate_bandwidth <- function(x, bw) {
if (identical(bw, "scott.diag")) {
h <- stats::sd(x) * length(x) ^ (-1 / 5)
if (!is.finite(h) || h <= 0) {
spread <- max(stats::sd(x), diff(range(x)), 1, na.rm = TRUE)
h <- spread / 10
}
return(h)
}
selector <- switch(
bw,
Hpi = stats::bw.SJ,
Hscv = stats::bw.ucv,
Hpi.diag = stats::bw.nrd0
)
h <- tryCatch(selector(x), error = function(e) NA_real_)
if (!is.finite(h) || h <= 0) {
h <- tryCatch(stats::density(x)$bw, error = function(e) NA_real_)
}
if (!is.finite(h) || h <= 0) {
spread <- max(stats::sd(x), diff(range(x)), 1, na.rm = TRUE)
h <- spread / 10
}
h
}
.scott_diag_bandwidth <- function(x) {
d <- ncol(x)
n <- nrow(x)
sds <- apply(x, 2, stats::sd)
bad <- !is.finite(sds) | sds <= 0
if (any(bad)) {
spreads <- apply(x[, bad, drop = FALSE], 2, function(col) {
max(stats::sd(col), diff(range(col)), 1, na.rm = TRUE)
})
sds[bad] <- spreads
}
h <- n ^ (-1 / (d + 4))
diag((h * sds) ^ 2, nrow = d, ncol = d)
}
.kde_1d_values <- function(x, eval_points, h, chunk_size = 5000L) {
out <- numeric(length(eval_points))
starts <- seq.int(1L, length(eval_points), by = chunk_size)
for (start in starts) {
stop <- min(start + chunk_size - 1L, length(eval_points))
z <- outer(eval_points[start:stop], x, `-`) / h
out[start:stop] <- rowMeans(stats::dnorm(z)) / h
}
out
}
.logsumexp <- function(x) {
m <- max(x)
if (!is.finite(m)) {
return(m)
}
m + log(sum(exp(x - m)))
}
.kde_diag_gaussian_values <- function(x, eval_points, H, chunk_size = 1000L) {
if (!is.matrix(H) || nrow(H) != ncol(H) || nrow(H) != ncol(x)) {
stop("Diagonal KDE engine received an invalid bandwidth matrix.", call. = FALSE)
}
off_diag <- H
diag(off_diag) <- 0
if (any(abs(off_diag) > sqrt(.Machine$double.eps))) {
stop(
"`engine = \"fast_diag\"` requires a diagonal bandwidth matrix. ",
"Use `bw = \"scott.diag\"` or `bw = \"Hpi.diag\"`.",
call. = FALSE
)
}
variances <- diag(H)
if (any(!is.finite(variances)) || any(variances <= 0)) {
stop("KDE bandwidth matrix must have positive finite diagonal entries.", call. = FALSE)
}
x <- as.matrix(x)
eval_points <- as.matrix(eval_points)
log_density <- numeric(nrow(eval_points))
inv_variances <- 1 / variances
starts <- seq.int(1L, nrow(eval_points), by = chunk_size)
for (start in starts) {
stop <- min(start + chunk_size - 1L, nrow(eval_points))
eval_chunk <- eval_points[start:stop, , drop = FALSE]
log_kernel <- matrix(0, nrow = nrow(eval_chunk), ncol = nrow(x))
for (j in seq_len(ncol(x))) {
diff <- outer(eval_chunk[, j], x[, j], `-`)
log_kernel <- log_kernel - 0.5 * diff * diff * inv_variances[j]
}
log_density[start:stop] <- apply(log_kernel, 1, .logsumexp) - log(nrow(x))
}
scale <- max(log_density)
if (!is.finite(scale)) {
return(rep(0, length(log_density)))
}
exp(log_density - scale)
}
.check_eval_seed <- function(eval_seed) {
if (is.null(eval_seed)) {
return(invisible(NULL))
}
if (!is.numeric(eval_seed) || length(eval_seed) != 1L ||
!is.finite(eval_seed) || eval_seed != trunc(eval_seed) ||
abs(eval_seed) > .Machine$integer.max) {
stop(
"`eval_seed` must be NULL or a single finite 32-bit integer.",
call. = FALSE
)
}
invisible(NULL)
}
# Deterministic Park-Miller uniforms for seeded internal draws. This generator
# is deliberately local: unlike set.seed(), it never reads or writes the
# caller's session RNG state. The Schrage update avoids integer overflow.
.local_uniforms <- function(n, eval_seed) {
.check_eval_seed(eval_seed)
.check_positive_count(n, "n")
if (is.null(eval_seed)) {
stop("`.local_uniforms()` requires a non-NULL seed.", call. = FALSE)
}
modulus <- 2147483647
multiplier <- 16807
quotient <- 127773
remainder <- 2836
state <- (as.double(eval_seed) %% (modulus - 1)) + 1
out <- numeric(as.integer(n))
for (i in seq_along(out)) {
hi <- floor(state / quotient)
lo <- state - hi * quotient
next_state <- multiplier * lo - remainder * hi
state <- if (next_state > 0) next_state else next_state + modulus
out[i] <- state / modulus
}
out
}
.local_normals <- function(n, eval_seed) {
.check_positive_count(n, "n")
n <- as.integer(n)
n_pairs <- ceiling(n / 2)
u <- .local_uniforms(2L * n_pairs, eval_seed)
radius <- sqrt(-2 * log(u[seq_len(n_pairs)]))
angle <- 2 * pi * u[n_pairs + seq_len(n_pairs)]
as.vector(rbind(radius * cos(angle), radius * sin(angle)))[seq_len(n)]
}
.local_sample_int <- function(n, size, eval_seed) {
.check_positive_count(n, "n")
.check_positive_count(size, "size")
n <- as.integer(n)
size <- as.integer(size)
if (size > n) {
stop("`size` must not exceed `n`.", call. = FALSE)
}
order(.local_uniforms(n, eval_seed))[seq_len(size)]
}
.sample_kde_eval_points <- function(eval_pts, eval_n = NULL, eval_seed = NULL) {
if (is.null(eval_n)) {
return(eval_pts)
}
.check_positive_count(eval_n, "eval_n")
eval_n <- as.integer(eval_n)
if (nrow(eval_pts) <= eval_n) {
return(eval_pts)
}
idx <- if (is.null(eval_seed)) {
sample.int(nrow(eval_pts), eval_n)
} else {
.local_sample_int(nrow(eval_pts), eval_n, eval_seed)
}
eval_pts[idx, , drop = FALSE]
}
.check_bw_scale <- function(bw_scale) {
if (!is.numeric(bw_scale) || length(bw_scale) != 1L ||
!is.finite(bw_scale) || bw_scale <= 0) {
stop("`bw_scale` must be a single positive finite number.", call. = FALSE)
}
invisible(bw_scale)
}
.scale_bandwidth <- function(bwspec, bw_scale = 1) {
# `bw_scale` multiplies the bandwidth on the standard-deviation scale: a
# univariate bandwidth h becomes bw_scale * h, and a bandwidth *matrix* H (a
# covariance-scale object) becomes bw_scale^2 * H. `bw_scale = 0.5` and `2`
# are the halved and doubled bandwidths of the smoothing-sensitivity check.
if (bw_scale == 1) {
return(bwspec)
}
if (is.matrix(bwspec)) bwspec * bw_scale^2 else bwspec * bw_scale
}
.select_multivariate_bandwidth <- function(x, bw, label) {
tryCatch(
switch(
bw,
Hpi = ks::Hpi(x),
Hscv = ks::Hscv(x),
Hpi.diag = ks::Hpi.diag(x),
scott.diag = .scott_diag_bandwidth(x)
),
error = function(e) {
stop(
"KDE bandwidth selection failed for category `", label, "`. ",
"Check that the category has more observations than feature dimensions ",
"and that feature columns are not constant or collinear. Original error: ",
conditionMessage(e),
call. = FALSE
)
}
)
}
.kde_density_pair <- function(data,
features,
category_col,
bw = c("Hpi", "Hscv", "Hpi.diag", "scott.diag"),
eval_on = c("pooled", "group1", "group2", "pooled_sample"),
eval_n = NULL,
eval_seed = NULL,
engine = c("ks", "fast_diag", "fast_diagonal"),
chunk_size = 1000L,
metric = "KDE",
bw_scale = 1) {
bw <- match.arg(bw)
eval_on <- match.arg(eval_on)
engine <- .match_kde_engine(engine)
.check_bw_scale(bw_scale)
if (!is.null(eval_n)) {
.check_positive_count(eval_n, "eval_n")
}
if (identical(eval_on, "pooled_sample") && is.null(eval_n)) {
stop("`eval_n` must be supplied when `eval_on = \"pooled_sample\"`.", call. = FALSE)
}
.check_positive_count(chunk_size, "chunk_size")
.check_columns(data, c(category_col, features))
data <- .metric_data(data, c(category_col, features))
.check_numeric_features(data, features)
levs <- .two_levels(data[[category_col]], "category_col")
n_features <- length(features)
.check_two_category_sample_size(
data,
category_col,
.kde_min_category_tokens(n_features),
metric
)
d1 <- data[data[[category_col]] == levs[1], , drop = FALSE]
d2 <- data[data[[category_col]] == levs[2], , drop = FALSE]
X1 <- as.matrix(d1[, features, drop = FALSE])
X2 <- as.matrix(d2[, features, drop = FALSE])
X_all <- as.matrix(data[, features, drop = FALSE])
eval_source <- if (identical(eval_on, "pooled_sample")) "pooled" else eval_on
eval_pts <- switch(
eval_source,
pooled = X_all,
group1 = X1,
group2 = X2
)
eval_pts <- .sample_kde_eval_points(eval_pts, eval_n = eval_n, eval_seed = eval_seed)
if (n_features == 1L) {
x1 <- as.numeric(X1[, 1])
x2 <- as.numeric(X2[, 1])
eval_vec <- as.numeric(eval_pts[, 1])
h1 <- .scale_bandwidth(.select_univariate_bandwidth(x1, bw), bw_scale)
h2 <- .scale_bandwidth(.select_univariate_bandwidth(x2, bw), bw_scale)
p <- .kde_1d_values(x1, eval_vec, h1)
q <- .kde_1d_values(x2, eval_vec, h2)
} else {
if (identical(engine, "fast_diag") &&
!bw %in% c("Hpi.diag", "scott.diag")) {
stop(
"`engine = \"fast_diag\"` requires `bw = \"scott.diag\"` or ",
"`bw = \"Hpi.diag\"` for multivariate KDE.",
call. = FALSE
)
}
H1 <- .scale_bandwidth(.select_multivariate_bandwidth(X1, bw, levs[1]), bw_scale)
H2 <- .scale_bandwidth(.select_multivariate_bandwidth(X2, bw, levs[2]), bw_scale)
if (identical(engine, "fast_diag")) {
p <- .kde_diag_gaussian_values(X1, eval_pts, H1, chunk_size = chunk_size)
q <- .kde_diag_gaussian_values(X2, eval_pts, H2, chunk_size = chunk_size)
} else {
kde1 <- tryCatch(
ks::kde(x = X1, H = H1, eval.points = eval_pts),
error = function(e) {
stop(
"KDE estimation failed for category `", levs[1], "`. Original error: ",
conditionMessage(e),
call. = FALSE
)
}
)
kde2 <- tryCatch(
ks::kde(x = X2, H = H2, eval.points = eval_pts),
error = function(e) {
stop(
"KDE estimation failed for category `", levs[2], "`. Original error: ",
conditionMessage(e),
call. = FALSE
)
}
)
p <- as.numeric(kde1$estimate)
q <- as.numeric(kde2$estimate)
}
}
if (any(!is.finite(p)) || any(!is.finite(q)) || sum(p) <= 0 || sum(q) <= 0) {
stop("KDE returned invalid density estimates.", call. = FALSE)
}
list(p = p, q = q, levels = levs, data = data)
}
# ---- Monte-Carlo plug-in KDE estimator ----------------------------------
# Consistent estimator of the continuous Jensen-Shannon divergence / overlap:
# evaluate each category's KDE at that category's own observations and average
# the true log density ratio against the mixture. Dimension-agnostic (unlike a
# grid) and unbiased in the limit (unlike the self-normalized sample-point plug
# -in used by `.kde_density_pair()`, kept as the `method = "legacy"` path).
.log_add_exp <- function(a, b) {
n <- max(length(a), length(b))
a <- rep(a, length.out = n)
b <- rep(b, length.out = n)
m <- pmax(a, b)
out <- m
# Only evaluate log1p where the max is finite; both -Inf stays -Inf. Indexing
# (not ifelse) avoids computing NaN intermediates that emit spurious warnings.
fin <- is.finite(m)
out[fin] <- m[fin] + log1p(exp(-abs(a[fin] - b[fin])))
out
}
.log_sub_exp <- function(a, b) {
# log(exp(a) - exp(b)); -Inf where a <= b (e.g., an isolated leave-one-out
# point). Evaluate the log only on strictly-greater elements so `log1p` is
# never handed a value <= -1 (which would emit a "NaNs produced" warning).
n <- max(length(a), length(b))
a <- rep(a, length.out = n)
b <- rep(b, length.out = n)
out <- rep(-Inf, n)
ok <- a > b
ok[is.na(ok)] <- FALSE
out[ok] <- a[ok] + log1p(-exp(b[ok] - a[ok]))
out
}
.kde_diag_log_density <- function(x, eval_points, H, chunk_size = 1000L) {
variances <- diag(H)
d <- ncol(x)
log_norm <- -0.5 * (d * log(2 * pi) + sum(log(variances)))
x <- as.matrix(x)
eval_points <- as.matrix(eval_points)
log_density <- numeric(nrow(eval_points))
inv_variances <- 1 / variances
starts <- seq.int(1L, nrow(eval_points), by = chunk_size)
for (start in starts) {
stop <- min(start + chunk_size - 1L, nrow(eval_points))
eval_chunk <- eval_points[start:stop, , drop = FALSE]
log_kernel <- matrix(0, nrow = nrow(eval_chunk), ncol = nrow(x))
for (j in seq_len(ncol(x))) {
diff <- outer(eval_chunk[, j], x[, j], `-`)
log_kernel <- log_kernel - 0.5 * diff * diff * inv_variances[j]
}
log_density[start:stop] <- apply(log_kernel, 1, .logsumexp) - log(nrow(x))
}
log_density + log_norm
}
.kde_kh0 <- function(bwspec, d) {
# K_H(0): the kernel's value at the origin (self-contribution), for LOO.
if (is.matrix(bwspec)) {
(2 * pi) ^ (-d / 2) * det(bwspec) ^ (-0.5)
} else {
stats::dnorm(0) / bwspec
}
}
.kde_eval_logdens <- function(train, eval, bwspec, engine, chunk_size, label) {
if (!is.matrix(bwspec)) {
dens <- .kde_1d_values(as.numeric(train[, 1]), as.numeric(eval[, 1]), bwspec)
return(log(dens))
}
if (identical(engine, "fast_diag")) {
return(.kde_diag_log_density(train, eval, bwspec, chunk_size = chunk_size))
}
kde <- tryCatch(
ks::kde(x = train, H = bwspec, eval.points = eval),
error = function(e) {
stop(
"KDE estimation failed for category `", label, "`. Original error: ",
conditionMessage(e), call. = FALSE
)
}
)
dens <- as.numeric(kde$estimate)
dens[dens < 0] <- 0
log(dens)
}
.select_kde_bandwidth <- function(train, bw, engine, n_features, label,
bw_scale = 1) {
if (n_features == 1L) {
return(.scale_bandwidth(
.select_univariate_bandwidth(as.numeric(train[, 1]), bw), bw_scale
))
}
if (identical(engine, "fast_diag") && !bw %in% c("Hpi.diag", "scott.diag")) {
stop(
"`engine = \"fast_diag\"` requires `bw = \"scott.diag\"` or ",
"`bw = \"Hpi.diag\"` for multivariate KDE.",
call. = FALSE
)
}
.scale_bandwidth(.select_multivariate_bandwidth(train, bw, label), bw_scale)
}
.kde_mc_pair <- function(data,
features,
category_col,
bw = c("Hpi", "Hscv", "Hpi.diag", "scott.diag"),
eval_n = NULL,
eval_seed = NULL,
engine = c("ks", "fast_diag", "fast_diagonal"),
chunk_size = 1000L,
metric = "KDE",
bw_scale = 1) {
bw <- match.arg(bw)
engine <- .match_kde_engine(engine)
.check_bw_scale(bw_scale)
if (!is.null(eval_n)) {
.check_positive_count(eval_n, "eval_n")
}
.check_positive_count(chunk_size, "chunk_size")
.check_columns(data, c(category_col, features))
data <- .metric_data(data, c(category_col, features))
.check_numeric_features(data, features)
levs <- .two_levels(data[[category_col]], "category_col")
n_features <- length(features)
.check_two_category_sample_size(
data, category_col, .kde_min_category_tokens(n_features), metric
)
X1 <- as.matrix(data[data[[category_col]] == levs[1], features, drop = FALSE])
X2 <- as.matrix(data[data[[category_col]] == levs[2], features, drop = FALSE])
n1 <- nrow(X1)
n2 <- nrow(X2)
# KDEs are trained on the full samples; evaluation points may be subsampled
# for speed (leave-one-out below still uses the full training size n1/n2).
X1e <- .sample_kde_eval_points(X1, eval_n = eval_n, eval_seed = eval_seed)
X2e <- .sample_kde_eval_points(X2, eval_n = eval_n, eval_seed = eval_seed)
bw1 <- .select_kde_bandwidth(X1, bw, engine, n_features, levs[1], bw_scale)
bw2 <- .select_kde_bandwidth(X2, bw, engine, n_features, levs[2], bw_scale)
out <- list(
logp1 = .kde_eval_logdens(X1, X1e, bw1, engine, chunk_size, levs[1]),
logq1 = .kde_eval_logdens(X2, X1e, bw2, engine, chunk_size, levs[2]),
logp2 = .kde_eval_logdens(X1, X2e, bw1, engine, chunk_size, levs[1]),
logq2 = .kde_eval_logdens(X2, X2e, bw2, engine, chunk_size, levs[2]),
n1 = n1, n2 = n2,
kh0_1 = .kde_kh0(bw1, n_features),
kh0_2 = .kde_kh0(bw2, n_features),
levels = levs, data = data
)
if (any(!is.finite(out$logp1)) && any(!is.finite(out$logp2))) {
stop("KDE returned invalid density estimates.", call. = FALSE)
}
out
}
.loo_alpha <- function(n) {
# Strength of the partial leave-one-out correction (see `.loo_logdens()`),
# phased in with sample size: alpha = 1/2 at n = 20 (the package's
# `min_tokens` default) and alpha -> 1 (full leave-one-out) as n grows.
n / (n + 20)
}
.loo_logdens <- function(log_dens, n, kh0, alpha = 1) {
# Partial leave-one-out log density at a KDE's own training points:
# p_alpha(x_i) = (n * p_hat(x_i) - alpha * K_H(0)) / (n - alpha),
# which removes a fraction `alpha` of the point's own kernel. `alpha = 1` is
# the classical leave-one-out density.
#
# Why partial: at an isolated point the self-kernel K_H(0)/n dominates the
# full-sample estimate p_hat(x_i), so full leave-one-out drives p_LOO toward
# 0 there. In the Jensen-Shannon integrand log2(p / m) those points produce
# unbounded *negative* contributions that drag the plug-in mean below its
# true (provably non-negative) value -- small real divergences were floored
# to exactly 0 by the final clamp. Removing only a fraction alpha < 1 keeps
# the density bounded below by (1 - alpha) * K_H(0) / (n - alpha) > 0, which
# bounds the integrand and eliminates the flooring, while still correcting
# the bulk of the resubstitution bias. With alpha = .loo_alpha(n) the
# correction approaches the full leave-one-out density as n grows, so the
# estimator keeps its consistency for the continuous Jensen-Shannon
# divergence; the retained self-kernel share vanishes with the same order as
# the KDE's own smoothing bias.
.log_sub_exp(log(n) + log_dens, log(alpha) + log(kh0)) - log(n - alpha)
}
.jsd_mc <- function(mc, loo = TRUE) {
ln2 <- log(2)
logp1 <- if (isTRUE(loo)) {
.loo_logdens(mc$logp1, mc$n1, mc$kh0_1, .loo_alpha(mc$n1))
} else {
mc$logp1
}
logm1 <- log(0.5) + .log_add_exp(logp1, mc$logq1)
t1 <- (logp1 - logm1) / ln2
logq2 <- if (isTRUE(loo)) {
.loo_logdens(mc$logq2, mc$n2, mc$kh0_2, .loo_alpha(mc$n2))
} else {
mc$logq2
}
logm2 <- log(0.5) + .log_add_exp(mc$logp2, logq2)
t2 <- (logq2 - logm2) / ln2
t1 <- t1[is.finite(t1)]
t2 <- t2[is.finite(t2)]
if (!length(t1) || !length(t2)) {
stop("Monte-Carlo JSD: no usable evaluation points.", call. = FALSE)
}
min(max(0.5 * mean(t1) + 0.5 * mean(t2), 0), 1)
}
.overlap_mc <- function(mc) {
# OVL = integral of min(p, q); estimate each half with that group's own
# samples via min(1, cross-density / self-density).
o1 <- pmin(1, exp(mc$logq1 - mc$logp1))
o2 <- pmin(1, exp(mc$logp2 - mc$logq2))
o1 <- o1[is.finite(o1)]
o2 <- o2[is.finite(o2)]
if (!length(o1) || !length(o2)) {
stop("Monte-Carlo overlap: no usable evaluation points.", call. = FALSE)
}
min(max(0.5 * mean(o1) + 0.5 * mean(o2), 0), 1)
}
.bhatt_mc <- function(mc, loo = TRUE) {
# Bhattacharyya coefficient BC = integral of sqrt(p q), read off the same
# density pair as JSD and overlap: E_P[sqrt(q / p)] from P's own samples and
# E_Q[sqrt(p / q)] from Q's, averaged. The self-densities take the same
# partial leave-one-out correction as `.jsd_mc()`, so the matched-kernel
# Bhattacharyya and Jensen-Shannon estimates share one estimator.
logp1 <- if (isTRUE(loo)) {
.loo_logdens(mc$logp1, mc$n1, mc$kh0_1, .loo_alpha(mc$n1))
} else {
mc$logp1
}
logq2 <- if (isTRUE(loo)) {
.loo_logdens(mc$logq2, mc$n2, mc$kh0_2, .loo_alpha(mc$n2))
} else {
mc$logq2
}
b1 <- exp(0.5 * (mc$logq1 - logp1))
b2 <- exp(0.5 * (mc$logp2 - logq2))
b1 <- b1[is.finite(b1)]
b2 <- b2[is.finite(b2)]
if (!length(b1) || !length(b2)) {
stop("Monte-Carlo Bhattacharyya: no usable evaluation points.", call. = FALSE)
}
min(max(0.5 * mean(b1) + 0.5 * mean(b2), 0), 1)
}
# ---- Parametric multivariate-normal density backend ---------------------
# `density = "mvnorm"`: fit one Gaussian per category and estimate JSD / overlap
# with the same Monte-Carlo plug-in used for KDE, but with the parametric
# density in place of the kernel estimate. A Gaussian fit has no self-kernel, so
# no leave-one-out correction is needed (`.jsd_mc(loo = FALSE)`). Jensen-Shannon
# divergence between two Gaussians has no closed form (the mixture is a Gaussian
# mixture), so this remains a Monte-Carlo estimate.
.mvn_logdens <- function(X, mu, S) {
# log N(x; mu, S) per row of X, via a Cholesky solve (no extra dependency).
d <- ncol(X)
R <- chol(S) # S = t(R) %*% R, R upper-triangular
logdet <- 2 * sum(log(diag(R)))
Xc <- sweep(as.matrix(X), 2, mu) # centered rows
z <- forwardsolve(t(R), t(Xc)) # t(R) lower-tri; z = t(R)^-1 (x - mu)
quad <- colSums(z * z) # (x - mu)' S^-1 (x - mu)
-0.5 * (d * log(2 * pi) + logdet + quad)
}
.mvn_rsample <- function(n, mu, S, standard_normals = NULL) {
# n independent draws from N(mu, S). With Z an n-by-d matrix of iid N(0, 1)
# entries and R = chol(S) (upper-triangular, S = t(R) %*% R), the rows of
# Z %*% R have covariance t(R) %*% R = S; adding mu shifts the mean.
d <- length(mu)
if (is.null(standard_normals)) {
standard_normals <- stats::rnorm(n * d)
}
if (length(standard_normals) != n * d || any(!is.finite(standard_normals))) {
stop(
"`standard_normals` must contain `n * length(mu)` finite values.",
call. = FALSE
)
}
Z <- matrix(standard_normals, nrow = n, ncol = d)
sweep(Z %*% chol(S), 2, mu, "+")
}
.mvn_fit_category <- function(X, ridge) {
# The single Gaussian fit used by the mvnorm density backend -- and by
# plot_contrast(), so plotted regions come from the same model as the metric.
list(mu = colMeans(X), S = stats::cov(X) + diag(ridge, ncol(X)))
}
.mvnorm_mc_pair <- function(data,
features,
category_col,
mc_n = 10000L,
eval_seed = NULL,
ridge = 1e-6,
metric = "mvnorm") {
.check_ridge_eps(ridge, "ridge")
.check_columns(data, c(category_col, features))
data <- .metric_data(data, c(category_col, features))
.check_numeric_features(data, features)
.check_positive_count(mc_n, "mc_n")
mc_n <- as.integer(mc_n)
levs <- .two_levels(data[[category_col]], "category_col")
d <- length(features)
.check_two_category_sample_size(
data, category_col, .kde_min_category_tokens(d), metric
)
X1 <- as.matrix(data[data[[category_col]] == levs[1], features, drop = FALSE])
X2 <- as.matrix(data[data[[category_col]] == levs[2], features, drop = FALSE])
fit1 <- .mvn_fit_category(X1, ridge)
fit2 <- .mvn_fit_category(X2, ridge)
mu1 <- fit1$mu; S1 <- fit1$S
mu2 <- fit2$mu; S2 <- fit2$S
if (!isTRUE(tryCatch({ chol(S1); chol(S2); TRUE }, error = function(e) FALSE))) {
stop(
"MVN density backend: a category covariance is not positive definite. ",
"Try increasing `ridge` or reducing feature dimensionality.",
call. = FALSE
)
}
# Fresh-sample estimator: draw mc_n points from each fitted Gaussian and
# evaluate both densities there. Unlike reusing the training points, this
# targets the JSD / overlap between the two *fitted* Gaussians -- a
# well-defined estimand independent of the observed sample -- with lower
# variance and no resubstitution bias, so no leave-one-out term is required.
if (is.null(eval_seed)) {
draws <- list(
X1e = .mvn_rsample(mc_n, mu1, S1),
X2e = .mvn_rsample(mc_n, mu2, S2)
)
} else {
standard_normals <- .local_normals(2L * mc_n * d, eval_seed)
split_at <- mc_n * d
draws <- list(
X1e = .mvn_rsample(
mc_n, mu1, S1, standard_normals[seq_len(split_at)]
),
X2e = .mvn_rsample(
mc_n, mu2, S2, standard_normals[split_at + seq_len(split_at)]
)
)
}
list(
logp1 = .mvn_logdens(draws$X1e, mu1, S1),
logq1 = .mvn_logdens(draws$X1e, mu2, S2),
logp2 = .mvn_logdens(draws$X2e, mu1, S1),
logq2 = .mvn_logdens(draws$X2e, mu2, S2),
n1 = nrow(X1), n2 = nrow(X2), levels = levs, data = data
)
}
Any scripts or data that you put into this service are public.
Add the following code to your website.
For more information on customizing the embed code, read Embedding Snippets.