Nothing
#' hBayesDM Model Base Function
#'
#' @description
#' The base function from which all hBayesDM model functions are created.
#'
#' Contributor: \href{https://ccs-lab.github.io/team/jethro-lee/}{Jethro Lee}
#'
#' @keywords internal
#'
#' @include settings.R
#' @include fit_cmdstan.R
#' @importFrom utils head
#' @importFrom stats complete.cases qnorm median
#' @importFrom data.table fread
#' @importFrom parallel detectCores
#'
#' @param task_name Character value for name of task. E.g. \code{"gng"}.
#' @param model_name Character value for name of model. E.g. \code{"m1"}.
#' @param model_type Character value for modeling type: \code{""} OR \code{"single"} OR
#' \code{"multipleB"}.
#' @param data_columns Character vector of necessary column names for the data. E.g.
#' \code{c("subjID", "cue", "keyPressed", "outcome")}.
#' @param parameters List of parameters, with information about their lower bound, plausible value,
#' upper bound. E.g. \code{list("xi" = c(0, 0.1, 1), "ep" = c(0, 0.2, 1), "rho" = c(0, exp(2),
#' Inf))}.
#' @param regressors List of regressors, with information about their extracted dimensions. E.g.
#' \code{list("Qgo" = 2, "Qnogo" = 2, "Wgo" = 2, "Wnogo" = 2)}. OR if model-based regressors are
#' not available for this model, \code{NULL}.
#' @param postpreds Character vector of name(s) for the trial-level posterior predictive
#' simulations. Default is \code{"y_pred"}. OR if posterior predictions are not yet available for
#' this model, \code{NULL}.
#' @param stanmodel_arg Leave as \code{NULL} (default) for completed models. Else should be either a
#' character value (the path to a Stan file) or a pre-compiled \code{cmdstanr::CmdStanModel}
#' object.
#' @param preprocess_func Function to preprocess the raw data before it gets passed to Stan. Takes
#' (at least) two arguments: a data.table object \code{raw_data} and a list object
#' \code{general_info}. Possible to include additional argument(s) to use during preprocessing.
#' Should return a list object \code{data_list}, which will then directly be passed to Stan.
#'
#' @details
#' \strong{task_name}: Typically same task models share the same data column requirements.
#'
#' \strong{model_name}: Typically different models are distinguished by their different list of
#' parameters.
#'
#' \strong{model_type} is one of the following three:
#' \describe{
#' \item{\code{""}}{Modeling of multiple subjects. (Default hierarchical Bayesian analysis.)}
#' \item{\code{"single"}}{Modeling of a single subject.}
#' \item{\code{"multipleB"}}{Modeling of multiple subjects, where multiple blocks exist within
#' each subject.}
#' }
#'
#' \strong{data_columns} must be the entirety of necessary data columns used at some point in the R
#' or Stan code. I.e. \code{"subjID"} must always be included. In the case of 'multipleB' type
#' models, \code{"block"} should also be included as well.
#'
#' \strong{parameters} is a list object, whose keys are the parameters of this model. Each parameter
#' key must be assigned a numeric vector holding 3 elements: the parameter's lower bound,
#' plausible value, and upper bound.
#'
#' \strong{regressors} is a list object, whose keys are the model-based regressors of this model.
#' Each regressor key must be assigned a numeric value indicating the number of dimensions its
#' data will be extracted as. If model-based regressors are not available for this model, this
#' argument should just be \code{NULL}.
#'
#' \strong{postpreds} defaults to \code{"y_pred"}, but any other character vector holding
#' appropriate names is possible (c.f. Two-Step Task models). If posterior predictions are not yet
#' available for this model, this argument should just be \code{NULL}.
#'
#' \strong{stanmodel_arg} can be used by developers, during the developmental stage of creating a
#' new model function. If this argument is passed a character value, the Stan file with the
#' corresponding name will be used for model fitting. If this argument is passed a
#' \code{stanmodel} object, that \code{stanmodel} object will be used for model fitting. When
#' creation of the model function is complete, this argument should just be left as \code{NULL}.
#'
#' \strong{preprocess_func} is the part of the code that is specific to the model, and is thus
#' written in the specific model R file.\cr
#' Arguments for this function are:
#' \describe{
#' \item{\code{raw_data}}{A data.table that holds the raw user data, which was read by using
#' \code{\link[data.table]{fread}}.}
#' \item{\code{general_info}}{A list that holds the general informations about the raw data, i.e.
#' \code{subjs}, \code{n_subj}, \code{t_subjs}, \code{t_max}, \code{b_subjs}, \code{b_max}.}
#' \item{\code{...}}{Optional additional argument(s) that specific model functions may want to
#' include. Examples of such additional arguments currently being used in hBayesDM models are:
#' \code{RTbound} (choiceRT_ddm models), \code{payscale} (igt models), and \code{trans_prob} (ts
#' models).}
#' }
#' Return value for this function should be:
#' \describe{
#' \item{\code{data_list}}{A list with appropriately named keys (as required by the model Stan
#' file), holding the fully preprocessed user data.}
#' }
#' NOTE: Syntax for data.table slightly differs from that of data.frame. If you want to use
#' \code{raw_data} as a data.frame when writing the \code{preprocess_func}, simply begin with the
#' line: \code{raw_data <- as.data.frame(raw_data)}.\cr
#' NOTE: Because of allowing case & underscore insensitive column names in user data,
#' \code{raw_data} columns must now be referenced by their lowercase non-underscored versions,
#' e.g. \code{"subjid"}, within the code of the preprocess function.\cr
#'
#' @return A specific hBayesDM model function.
hBayesDM_model <- function(task_name = "",
model_name,
model_type = "",
data_columns,
parameters,
additional_args = NULL,
regressors = NULL,
postpreds = "y_pred",
stanmodel_arg = NULL,
preprocess_func) {
# The resulting hBayesDM model function to be returned
function(data = NULL,
niter = 4000,
nwarmup = 1000,
nchain = 4,
ncore = 1,
nthin = 1,
inits = "vb",
ind_pars = "mean",
model_regressor = FALSE,
vb = FALSE,
inc_postpred = FALSE,
adapt_delta = 0.95,
stepsize = 1,
max_treedepth = 10,
seed = 42,
...) {
############### Stop checks ###############
# Check if regressor available for this model
if (model_regressor && is.null(regressors)) {
stop("** Model-based regressors are not available for this model. **\n")
}
# Check if postpred available for this model
if (inc_postpred && is.null(postpreds)) {
stop("** Posterior predictions are not yet available for this model. **\n")
}
if (is.data.frame(data)) {
# Use the given data object
raw_data <- data.table::as.data.table(data)
} else if (length(data) == 1 && is.character(data)) {
# Set
if (data == "example") {
model_meta <- c()
if (task_name != "") {
model_meta <- c(model_meta, task_name)
} else {
model_meta <- c(model_meta, model_name)
}
if (model_type != "") {
model_meta <- c(model_meta, model_type)
}
if (length(model_meta) == 0) {
stop("invalid model configuration")
}
example_data <- paste0(paste(model_meta, collapse = "_"), "_exampleData.txt")
datafile <- system.file("extdata", example_data, package = "hBayesDM")
if (!file.exists(datafile)) {
stop("** Example data for this task does not exist **")
}
} else if (data == "choose") {
datafile <- file.choose()
} else {
datafile <- data
}
# Check if data file exists
if (!file.exists(datafile)) {
stop("** Data file does not exist. Please check again. **\n",
" e.g. data = \"MySubFolder/myData.txt\"\n")
}
# Load the data
raw_data <- data.table::fread(file = datafile, header = TRUE, sep = "\t", data.table = TRUE,
fill = TRUE, stringsAsFactors = TRUE, logical01 = FALSE)
# NOTE: Separator is fixed to "\t" because fread() has trouble reading space delimited files
# that have missing values.
} else {
stop("Invalid input for the 'data' value. ",
"You should pass a data.frame, or a filepath for a TSV file,",
"\"example\" for an example dataset, ",
"or \"choose\" to choose it in a prompt.")
}
# Save initial colnames of raw_data for later
colnames_raw_data <- colnames(raw_data)
# Check if necessary data columns all exist (while ignoring case and underscores)
insensitive_data_columns <- tolower(gsub("_", "", data_columns, fixed = TRUE))
colnames(raw_data) <- tolower(gsub("_", "", colnames(raw_data), fixed = TRUE))
if (!all(insensitive_data_columns %in% colnames(raw_data))) {
stop("** Data file is missing one or more necessary data columns. Please check again. **\n",
" Necessary data columns are: \"", paste0(data_columns, collapse = "\", \""), "\".\n")
}
# Remove only the rows containing NAs in necessary columns
complete_rows <- complete.cases(raw_data[, insensitive_data_columns, with = FALSE])
sum_incomplete_rows <- sum(!complete_rows)
if (sum_incomplete_rows > 0) {
raw_data <- raw_data[complete_rows, ]
cat("\n")
cat("The following lines of the data file have NAs in necessary columns:\n")
cat(paste0(head(which(!complete_rows), 100) + 1, collapse = ", "))
if (sum_incomplete_rows > 100) {
cat(", ...")
}
cat(" (total", sum_incomplete_rows, "lines)\n")
cat("These rows are removed prior to modeling the data.\n")
}
####################################################
## Prepare general info about the raw data #####
####################################################
subjs <- NULL # List of unique subjects (1D)
n_subj <- NULL # Total number of subjects (0D)
b_subjs <- NULL # Number of blocks per each subject (1D)
b_max <- NULL # Maximum number of blocks across all subjects (0D)
t_subjs <- NULL # Number of trials (per block) per subject (2D or 1D)
t_max <- NULL # Maximum number of trials across all blocks & subjects (0D)
# To avoid NOTEs by R CMD check
.N <- NULL
subjid <- NULL
if ((model_type == "") || (model_type == "single")) {
DT_trials <- raw_data[, .N, by = "subjid"]
subjs <- DT_trials$subjid
n_subj <- length(subjs)
t_subjs <- DT_trials$N
t_max <- max(t_subjs)
if ((model_type == "single") && (n_subj != 1)) {
stop("** More than 1 unique subjects exist in data file,",
" while using 'single' type model. **\n")
}
} else { # (model_type == "multipleB")
DT_trials <- raw_data[, .N, by = c("subjid", "block")]
DT_blocks <- DT_trials[, .N, by = "subjid"]
subjs <- DT_blocks$subjid
n_subj <- length(subjs)
b_subjs <- DT_blocks$N
b_max <- max(b_subjs)
t_subjs <- array(0, c(n_subj, b_max))
for (i in 1:n_subj) {
subj <- subjs[i]
b <- b_subjs[i]
t_subjs[i, 1:b] <- DT_trials[subjid == subj]$N
}
t_max <- max(t_subjs)
}
general_info <- list(subjs, n_subj, b_subjs, b_max, t_subjs, t_max)
names(general_info) <- c("subjs", "n_subj", "b_subjs", "b_max", "t_subjs", "t_max")
#########################################################
## Prepare: data_list #####
## pars #####
## model_name #####
#########################################################
# Preprocess the raw data to pass to Stan
if (is.null(additional_args)) {
data_list <- preprocess_func(raw_data, general_info, ...)
} else {
args <- list(...)
# set default values if not specified in args
for (nm in names(additional_args)) {
if (!nm %in% names(args)) {
# Single-bracket list assignment so NULL defaults survive (the
# double-bracket form would drop the entry instead).
args[nm] <- list(additional_args[[nm]])
}
}
data_list <- do.call(preprocess_func, c(list(raw_data, general_info), args))
}
# The parameters of interest for Stan
pars <- character()
if (model_type != "single") {
pars <- c(pars, paste0("mu_", names(parameters)), "sigma")
}
pars <- c(pars, names(parameters))
if ((task_name == "dd") && (model_type == "single")) {
log_parameter1 <- paste0("log", toupper(names(parameters)[1]))
pars <- c(pars, log_parameter1)
}
if ((model_name == "hgf_ibrb") && (model_type == "single")) {
pars <- c(pars, paste0("logit_", names(parameters)))
}
pars <- c(pars, "log_lik")
if (model_regressor) {
pars <- c(pars, names(regressors))
}
if (inc_postpred) {
pars <- c(pars, postpreds)
}
# Full name of model
model_meta <- c()
if (task_name != "") {
model_meta <- c(model_meta, task_name)
}
if (model_name != "") {
model_meta <- c(model_meta, model_name)
}
if (model_type != "") {
model_meta <- c(model_meta, model_type)
}
if (length(model_meta) == 0) {
stop("invalid model configuration")
}
model <- paste(model_meta, collapse = "_")
# Set number of cores for parallel computing
if (ncore <= 1) {
ncore <- 1
} else {
local_cores <- parallel::detectCores()
if (ncore > local_cores) {
ncore <- local_cores
warning("Number of cores specified for parallel computing greater than",
" number of locally available cores. Using all locally available cores.\n")
}
}
options(mc.cores = ncore)
############### Print for user ###############
cat("\n")
cat("Model name =", model, "\n")
if (is.character(data))
cat("Data file =", data, "\n")
cat("\n")
cat("Details:\n")
if (vb) {
cat(" Using variational inference\n")
} else {
cat(" # of chains =", nchain, "\n")
cat(" # of cores used =", ncore, "\n")
cat(" # of MCMC samples (per chain) =", niter, "\n")
cat(" # of burn-in samples =", nwarmup, "\n")
}
cat(" # of subjects =", n_subj, "\n")
if (model_type == "multipleB") {
cat(" # of (max) blocks per subject =", b_max, "\n")
}
if (model_type == "") {
cat(" # of (max) trials per subject =", t_max, "\n")
} else if (model_type == "multipleB") {
cat(" # of (max) trials...\n")
cat(" ...per block per subject =", t_max, "\n")
} else {
cat(" # of trials (for this subject) =", t_max, "\n")
}
# Models with additional arguments
if ((task_name == "choiceRT") && (model_name == "ddm")) {
RTbound <- list(...)$RTbound
cat(" `RTbound` is set to =", ifelse(is.null(RTbound), 0.1, RTbound), "\n")
}
if (task_name == "igt") {
payscale <- list(...)$payscale
cat(" `payscale` is set to =", ifelse(is.null(payscale), 100, payscale), "\n")
}
if (task_name == "ts") {
trans_prob <- list(...)$trans_prob
cat(" `trans_prob` is set to =", ifelse(is.null(trans_prob), 0.7, trans_prob), "\n")
}
# When extracting model-based regressors
if (model_regressor) {
cat("\n")
cat("**************************************\n")
cat("** Extract model-based regressors **\n")
cat("**************************************\n")
}
# An empty newline before Stan begins
if (nchain > 1) {
cat("\n")
}
# The Stan model name to use (cmdstanr will lazily compile on first use)
stan_model_name <- if (is.null(stanmodel_arg)) model else stanmodel_arg
# Initial values for the parameters
gen_init <- NULL
if (inits[1] == "vb") {
if (vb) {
cat("\n")
cat("*****************************************\n")
cat("** Use random values as initial values **\n")
cat("*****************************************\n")
gen_init <- "random"
} else {
cat("\n")
cat("****************************************\n")
cat("** Use VB estimates as initial values **\n")
cat("****************************************\n")
make_gen_init_from_vb <- function() {
stan_model_obj <- .hbayesdm_compile(stan_model_name)
fit_vb <- stan_model_obj$variational(data = data_list)
draws_df <- posterior::as_draws_df(fit_vb$draws())
m_vb <- colMeans(as.data.frame(draws_df))
function() {
ret <- list(
mu_pr = as.vector(m_vb[startsWith(names(m_vb), "mu_pr")]),
sigma = as.vector(m_vb[startsWith(names(m_vb), "sigma")])
)
for (p in names(parameters)) {
ret[[paste0(p, "_pr")]] <-
as.vector(m_vb[startsWith(names(m_vb), paste0(p, "_pr"))])
}
return(ret)
}
}
gen_init <- tryCatch(make_gen_init_from_vb(), error = function(e) {
cat("\n")
cat("******************************************\n")
cat("** Failed to obtain VB estimates. **\n")
cat("** Use random values as initial values. **\n")
cat("******************************************\n")
return("random")
})
}
} else if (inits[1] == "random") {
cat("\n")
cat("*****************************************\n")
cat("** Use random values as initial values **\n")
cat("*****************************************\n")
gen_init <- "random"
} else if (inits == 0) {
gen_init <- 0
} else {
if (inits[1] == "fixed") {
# plausible values of each parameter
inits <- unlist(lapply(parameters, "[", 2))
} else if (length(inits) != length(parameters)) {
stop("** Length of 'inits' must be ", length(parameters), " ",
"(= the number of parameters of this model). ",
"Please check again. **\n")
}
if (model_type == "single") {
gen_init <- function() {
individual_level <- as.list(inits)
names(individual_level) <- names(parameters)
return(individual_level)
}
} else {
gen_init <- function() {
primes <- numeric(length(parameters))
for (i in 1:length(parameters)) {
lb <- parameters[[i]][1] # lower bound
ub <- parameters[[i]][3] # upper bound
if (is.infinite(lb)) {
primes[i] <- inits[i] # (-Inf, Inf)
} else if (is.infinite(ub)) {
primes[i] <- log(inits[i] - lb) # ( lb, Inf)
} else {
primes[i] <- qnorm((inits[i] - lb) / (ub - lb)) # ( lb, ub)
}
}
group_level <- list(mu_pr = primes,
sigma = rep(1.0, length(primes)))
individual_level <- lapply(primes, function(x) rep(x, n_subj))
names(individual_level) <- paste0(names(parameters), "_pr")
return(c(group_level, individual_level))
}
}
}
############### Fit & extract ###############
fit_result <- .hbayesdm_fit(
model_name = stan_model_name,
data_list = data_list,
pars = pars,
gen_init = gen_init,
vb = vb,
nchain = nchain,
niter = niter,
nwarmup = nwarmup,
nthin = nthin,
adapt_delta = adapt_delta,
stepsize = stepsize,
max_treedepth = max_treedepth,
ncore = ncore,
seed = seed,
inc_postpred = inc_postpred,
postpreds = postpreds
)
fit <- fit_result$fit
par_vals <- fit_result$par_vals
# Define measurement of individual parameters
measure_ind_pars <- switch(ind_pars, mean = mean, median = median, mode = estimate_mode)
# Define which individual parameters to measure
which_ind_pars <- names(parameters)
if ((task_name == "dd") && (model_type == "single")) {
which_ind_pars <- c(which_ind_pars, log_parameter1)
}
# Measure all individual parameters (per subject)
compute_individual_params <- function(x, i = NULL) {
a <- par_vals[[x]]
d <- dim(a)
if (model_type == "single") {
if (is.null(d) || length(d) == 1) {
val <- measure_ind_pars(a) # real typed parameter
names(val) <- x
return(val)
} else if (length(d) == 2) {
param_cnt <- ncol(a)
if (param_cnt == 0) return(numeric(0))
val <- apply(a, 2, measure_ind_pars) # vector typed parameter (multiple parameters for single subject)
names(val) <- paste0(x, "[", seq_along(val), "]")
return(val)
}
} else {
stopifnot(!is.null(i))
if (length(d) == 2) {
param_cnt <- d[2]
if (param_cnt == 0) return(numeric(0))
val <- measure_ind_pars(a[, i]) # vector typed parameter (one for each subject)
names(val) <- x
return(val)
} else if (length(d) == 3) {
param_cnt <- d[3]
if (param_cnt == 0) return(numeric(0))
slice <- a[, i, , drop = TRUE]
if (is.null(dim(slice))) slice <- cbind(slice)
vals <- apply(slice, 2L, measure_ind_pars)
names(vals) <- paste0(x, "[", seq_along(vals), "]")
return(vals)
}
}
stop(sprintf("Unexpected shape for %s", x))
}
if (model_type == "single") {
first_row_param <- unlist(lapply(which_ind_pars, compute_individual_params), use.names = TRUE)
all_ind_pars <- as.data.frame(t(first_row_param), check.names = FALSE)
all_ind_pars <- cbind(subjID = subjs[1], all_ind_pars, row.names = NULL)
} else {
first_subj_params <- unlist(lapply(which_ind_pars, function(x) compute_individual_params(x, i = 1)), use.names = TRUE)
rows <- lapply(seq_len(n_subj), function(i) {
unlist(lapply(which_ind_pars, function(x) compute_individual_params(x, i = i)), use.names = TRUE)
})
mat <- do.call(rbind, rows)
colnames(mat) <- names(first_subj_params)
all_ind_pars <- cbind(subjID = subjs, as.data.frame(mat, check.names = FALSE), row.names = NULL)
}
# Model regressors (for model-based neuroimaging, etc.)
regressor_list <- NULL
if (model_regressor) {
regressor_list <- list()
for (r in names(regressors)) {
regressor_list[[r]] <- apply(par_vals[[r]], c(1:regressors[[r]]) + 1, measure_ind_pars)
}
}
# Give back initial colnames and revert data.table to data.frame
colnames(raw_data) <- colnames_raw_data
raw_data <- as.data.frame(raw_data)
# Wrap up data into a list
model_data <- list()
model_data$model <- model
model_data$all_ind_pars <- all_ind_pars
model_data$par_vals <- par_vals
model_data$fit <- fit
model_data$raw_data <- raw_data
if (model_regressor) {
model_data$model_regressor <- regressor_list
}
# Object class definition
class(model_data) <- "hBayesDM"
# Inform user of completion
cat("\n")
cat("************************************\n")
cat("**** Model fitting is complete! ****\n")
cat("************************************\n")
return(model_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.