Skip to content

Instantly share code, notes, and snippets.

@VisruthSK
Created March 2, 2026 05:50
Show Gist options
  • Select an option

  • Save VisruthSK/57eb18774d5ddd10b8b10643c3111d13 to your computer and use it in GitHub Desktop.

Select an option

Save VisruthSK/57eb18774d5ddd10b8b10643c3111d13 to your computer and use it in GitHub Desktop.
Mirai tests
project_root <- normalizePath(getwd(), mustWork = TRUE)
script_baseline <- file.path(project_root, "loo_mirai", "test.R")
script_fast <- file.path(project_root, "loo_mirai", "test_fast.R")
cache_baseline <- file.path(project_root, "loo_mirai", ".cache_baseline")
cache_fast <- file.path(project_root, "loo_mirai", ".cache_fast")
resolve_rscript <- function() {
candidate <- file.path(
R.home("bin"),
if (.Platform$OS.type == "windows") "Rscript.exe" else "Rscript"
)
if (file.exists(candidate)) {
return(candidate)
}
candidate <- Sys.which("Rscript")
if (!nzchar(candidate)) {
stop("Could not find Rscript in this R installation or PATH.")
}
candidate
}
extract_vector_lines <- function(lines) {
grep("^\\s*\\[[0-9]+\\]", lines, value = TRUE)
}
run_script <- function(rscript_bin, script_path, clear_dirs = character()) {
for (d in clear_dirs) {
unlink(d, recursive = TRUE, force = TRUE)
}
elapsed <- system.time(
output <- suppressWarnings(system2(
rscript_bin,
script_path,
stdout = TRUE,
stderr = TRUE
))
)[["elapsed"]]
status <- attr(output, "status")
if (is.null(status)) {
status <- 0L
}
list(
elapsed_sec = as.numeric(elapsed),
status = as.integer(status),
output = output,
vec_lines = extract_vector_lines(output)
)
}
summarize_run <- function(label, run) {
cat(sprintf("%s: %.3f sec (status=%d)\n", label, run$elapsed_sec, run$status))
if (run$status != 0) {
cat(sprintf("--- %s last output lines ---\n", label))
cat(paste(tail(run$output, 40), collapse = "\n"), "\n")
}
}
run_case <- function(label, rscript_bin, clear = FALSE) {
cat(sprintf("\n=== %s Runs ===\n", label))
baseline <- run_script(
rscript_bin,
script_baseline,
clear_dirs = if (clear) cache_baseline else character()
)
fast <- run_script(
rscript_bin,
script_fast,
clear_dirs = if (clear) cache_fast else character()
)
summarize_run(sprintf("%s_baseline", tolower(label)), baseline)
summarize_run(sprintf("%s_fast", tolower(label)), fast)
if (baseline$status != 0 || fast$status != 0) {
stop(sprintf("%s run failed; inspect output above.", label))
}
ok <- identical(baseline$vec_lines, fast$vec_lines)
cat(sprintf(
"%s_values_equal: %s\n",
tolower(label),
if (ok) "PASS" else "FAIL"
))
cat(sprintf(
"%s_speedup_x: %.3f\n",
tolower(label),
baseline$elapsed_sec / fast$elapsed_sec
))
ok
}
rscript_bin <- resolve_rscript()
ok_cold <- run_case("Cold", rscript_bin, clear = TRUE)
ok_warm <- run_case("Warm", rscript_bin, clear = FALSE)
if (!ok_cold || !ok_warm) {
stop("Baseline and fast scripts did not return the same values.")
}
# === Cold Runs ===
# cold_baseline: 135.110 sec (status=0)
# cold_fast: 85.700 sec (status=0)
# cold_values_equal: PASS
# cold_speedup_x: 1.577
# === Warm Runs ===
# warm_baseline: 130.830 sec (status=0)
# warm_fast: 64.340 sec (status=0)
# warm_values_equal: PASS
# warm_speedup_x: 2.033
library(cmdstanr)
data_bin <- list(N = 10, y = c(rep(1, 9), 0))
code_binom <- normalizePath(
file.path(cmdstan_path(), "examples", "bernoulli", "bernoulli.stan"),
mustWork = TRUE
)
build_root <- file.path(getwd(), "loo_mirai", ".cache_baseline")
dir.create(build_root, recursive = TRUE, showWarnings = FALSE)
host_dir <- file.path(build_root, "host")
dir.create(host_dir, recursive = TRUE, showWarnings = FALSE)
model_bin <- cmdstan_model(
stan_file = code_binom,
dir = host_dir,
compile_model_methods = TRUE,
force_recompile = TRUE
)
fit_bin <- model_bin$sample(data = data_bin, refresh = 0)
df_bin <- data.frame(theta = plogis(seq(-4, 6, length.out = 100))) |>
dplyr::mutate(pdfbeta = dbeta(theta, 9 + 1, 1 + 1))
# --- Parallel mirai_map ---
library(mirai)
mirai::daemons(n = 4)
# Get the absolute path to the Stan file so daemons can find it
stan_file_path <- normalizePath(code_binom)
# Initialize each daemon: load cmdstanr, compile model methods,
# fit the model, and store the fit in the daemon's global environment.
# The Stan file and data are small so passing them via ... is fine.
# everywhere() returns a mirai_map object (list of mirai, one per daemon).
# Collect with [] to block until all daemons finish initialization.
# Any daemon errors will surface here.
everywhere(
{
library(cmdstanr)
options(mc.cores = 1)
daemon_dir <- file.path(build_root, paste0("daemon_", Sys.getpid()))
dir.create(daemon_dir, recursive = TRUE, showWarnings = FALSE)
model_daemon <- cmdstan_model(
stan_file = stan_file_path,
dir = daemon_dir,
compile_model_methods = TRUE,
force_recompile = TRUE
)
fit_daemon <<- model_daemon$sample(
data = data_bin,
fixed_param = TRUE,
iter = 1,
iter_sampling = 1
)
fit_daemon$init_model_methods()
},
stan_file_path = stan_file_path,
build_root = build_root,
data_bin = data_bin,
.min = 4
)[]
# The mapped function uses 'fit_daemon' from the daemon's global environment
# — no external pointers need to be serialized
fit_pdf_daemon <- function(th) {
exp(fit_daemon$log_prob(
fit_daemon$unconstrain_variables(list(theta = th)),
jacobian = FALSE
))
}
result_mirai <- mirai_map(df_bin$theta, fit_pdf_daemon)[.flat]
print(result_mirai)
mirai::daemons(0)
library(cmdstanr)
library(mirai)
options(mc.cores = 1)
daemon_count <- 4
n_theta <- 100
stan_file_path <- normalizePath(
file.path(cmdstan_path(), "examples", "bernoulli", "bernoulli.stan"),
mustWork = TRUE
)
data_bin <- list(N = 10, y = c(rep(1, 9), 0))
theta <- plogis(seq(-4, 6, length.out = n_theta))
cache_root <- file.path(getwd(), "loo_mirai", ".cache_fast")
build_dir <- file.path(cache_root, "cmdstan")
rcpp_cache_dir <- file.path(cache_root, "rcpp")
dir.create(build_dir, recursive = TRUE, showWarnings = FALSE)
dir.create(rcpp_cache_dir, recursive = TRUE, showWarnings = FALSE)
options(rcpp.cache.dir = rcpp_cache_dir)
model <- cmdstan_model(
stan_file = stan_file_path,
dir = build_dir,
compile_model_methods = FALSE,
force_recompile = FALSE,
quiet = TRUE
)
exe_path <- normalizePath(model$exe_file(), mustWork = TRUE)
model_hpp <- cmdstan_model(
stan_file = stan_file_path,
exe_file = exe_path,
compile = FALSE
)
model_hpp$compile(force_recompile = TRUE, dry_run = TRUE, quiet = TRUE)
hpp_code <- model_hpp$.__enclos_env__$private$model_methods_env_$hpp_code_
warm_model <- cmdstan_model(
stan_file = stan_file_path,
exe_file = exe_path,
compile = FALSE
)
warm_fit <- warm_model$sample(
data = data_bin,
fixed_param = TRUE,
chains = 1,
iter_warmup = 0,
iter_sampling = 1,
refresh = 0
)
warm_fit$.__enclos_env__$private$model_methods_env_$hpp_code_ <- hpp_code
warm_fit$init_model_methods()
mirai::daemons(daemon_count)
on.exit(mirai::daemons(0), add = TRUE)
everywhere(
{
library(cmdstanr)
options(mc.cores = 1, rcpp.cache.dir = rcpp_cache_dir)
model_daemon <- cmdstan_model(
stan_file = stan_file_path,
exe_file = exe_path,
compile = FALSE
)
fit_daemon <<- model_daemon$sample(
data = data_bin,
fixed_param = TRUE,
chains = 1,
iter_warmup = 0,
iter_sampling = 1,
refresh = 0
)
fit_daemon$.__enclos_env__$private$model_methods_env_$hpp_code_ <- hpp_code
fit_daemon$init_model_methods()
},
stan_file_path = stan_file_path,
exe_path = exe_path,
data_bin = data_bin,
hpp_code = hpp_code,
rcpp_cache_dir = rcpp_cache_dir,
.min = daemon_count
)[]
fit_pdf_daemon <- function(th) {
exp(
fit_daemon$log_prob(
fit_daemon$unconstrain_variables(list(theta = th)),
jacobian = FALSE
)
)
}
result_mirai <- mirai::mirai_map(theta, fit_pdf_daemon)[.flat]
print(result_mirai)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment