Skip to content

Commit 9cd12fb

Browse files
committed
fix pathfinder column order in draws()
1 parent df60989 commit 9cd12fb

3 files changed

Lines changed: 35 additions & 12 deletions

File tree

R/csv.R

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -213,7 +213,13 @@ read_cmdstan_csv <- function(files,
213213
lp = lp
214214
))
215215
}
216-
user_variables_subset <- FALSE
216+
if (metadata$method == "pathfinder") {
217+
pathfinder_variables <- c("lp__", "lp_approx__", "path__")
218+
metadata$variables <- c(
219+
intersect(pathfinder_variables, metadata$variables),
220+
setdiff(metadata$variables, pathfinder_variables)
221+
)
222+
}
217223
if (is.null(variables)) { # variables = NULL returns all
218224
variables <- metadata$variables
219225
} else if (!any(nzchar(variables))) { # if variables = "" returns none
@@ -225,7 +231,6 @@ read_cmdstan_csv <- function(files,
225231
paste(res$not_found, collapse = ", "), call. = FALSE)
226232
}
227233
variables <- unrepair_variable_names(res$matching)
228-
user_variables_subset <- TRUE
229234
}
230235
if (is.null(sampler_diagnostics)) {
231236
sampler_diagnostics <- metadata$sampler_diagnostics
@@ -291,15 +296,6 @@ read_cmdstan_csv <- function(files,
291296
if (length(variables) > 0) {
292297
draws_list_id <- length(draws) + 1
293298
warmup_draws_list_id <- length(warmup_draws) + 1
294-
if (metadata$method == "pathfinder") {
295-
metadata$variables <- union(metadata$sampler_diagnostics, metadata$variables)
296-
if (!user_variables_subset) {
297-
# because for pathfinder variables and diagnostics are read in together,
298-
# if user hasn't selected a custom subset of variables we need to include
299-
# all diagnostics
300-
variables <- union(metadata$sampler_diagnostics, variables)
301-
}
302-
}
303299
suppressWarnings(
304300
draws[[draws_list_id]] <- data.table::fread(
305301
cmd = fread_cmd,
@@ -731,8 +727,12 @@ read_csv_metadata <- function(csv_file) {
731727
# if no # at the start of line, the line is the CSV header
732728
all_names <- strsplit(line, ",")[[1]]
733729
if (all(csv_file_info$algorithm != "fixed_param")) {
730+
non_sampler_diagnostics <- c("lp__", "log_p__", "log_g__", "log_q__")
731+
if (csv_file_info$method == "pathfinder") {
732+
non_sampler_diagnostics <- c(non_sampler_diagnostics, "lp_approx__", "path__")
733+
}
734734
csv_file_info[["sampler_diagnostics"]] <- all_names[endsWith(all_names, "__")]
735-
csv_file_info[["sampler_diagnostics"]] <- csv_file_info[["sampler_diagnostics"]][!(csv_file_info[["sampler_diagnostics"]] %in% c("lp__", "log_p__", "log_g__", "log_q__"))]
735+
csv_file_info[["sampler_diagnostics"]] <- csv_file_info[["sampler_diagnostics"]][!(csv_file_info[["sampler_diagnostics"]] %in% non_sampler_diagnostics)]
736736
csv_file_info[["variables"]] <- all_names[!(all_names %in% csv_file_info[["sampler_diagnostics"]])]
737737
} else {
738738
csv_file_info[["variables"]] <- all_names[!endsWith(all_names, "__")]

tests/testthat/test-csv.R

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -392,6 +392,25 @@ test_that("read_cmdstan_csv() works for laplace", {
392392
expect_equal(posterior::variables(csv_output_5$draws), c("alpha", "beta[2]"))
393393
})
394394

395+
test_that("read_cmdstan_csv() works for pathfinder", {
396+
csv_output <- read_cmdstan_csv(fit_logistic_pathfinder$output_files())
397+
expected_variables <- c(
398+
"lp__", "lp_approx__", "path__", "alpha",
399+
"beta[1]", "beta[2]", "beta[3]"
400+
)
401+
expect_equal(posterior::variables(csv_output$draws), expected_variables)
402+
expect_equal(csv_output$metadata$variables, expected_variables)
403+
404+
filtered_output <- read_cmdstan_csv(
405+
fit_logistic_pathfinder$output_files(),
406+
variables = c("path__", "lp_approx__")
407+
)
408+
expect_equal(
409+
posterior::variables(filtered_output$draws),
410+
c("path__", "lp_approx__")
411+
)
412+
})
413+
395414

396415
test_that("read_cmdstan_csv() works for generate_quantities", {
397416
csv_output_1 <- read_cmdstan_csv(fit_gq$output_files())

tests/testthat/test-model-pathfinder.R

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -103,6 +103,10 @@ expect_pathfinder_output <- function(object, num_chains = NULL) {
103103
test_that("Pathfinder Runs", {
104104
expect_pathfinder_output(fit <- mod$pathfinder(data=data_list, seed=1234, refresh = 0))
105105
expect_s3_class(fit, "CmdStanPathfinder")
106+
expect_equal(
107+
posterior::variables(fit$draws()),
108+
c("lp__", "lp_approx__", "path__", "theta")
109+
)
106110
})
107111

108112
test_that("pathfinder() method works with data files", {

0 commit comments

Comments
 (0)