Created
March 2, 2026 05:50
-
-
Save VisruthSK/57eb18774d5ddd10b8b10643c3111d13 to your computer and use it in GitHub Desktop.
Mirai tests
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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 |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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