Fits latent Dirichlet allocation (LDA), supervised topic models, and multilevel supervised topic models for text data with multiple outcome variables. Core estimation routines are implemented in C++ using the 'Rcpp' ecosystem. For topic models, see Blei et al. (2003) < https://www.jmlr.org/papers/volume3/blei03a/blei03a.pdf>. For supervised topic models, see Blei and McAuliffe (2007) < https://papers.nips.cc/paper_files/paper/2007/hash/d56b9fc4b0f1be8871f5e1c40c0067e7-Abstract.html>.
mlstm: Multilevel Supervised Topic Models with Multiple Outcomes in R
mlstm implements Multilevel Supervised Topic Models (MLSTM),
a probabilistic framework for analyzing text data with multiple
associated outcome variables.
Unlike standard supervised topic models that assume a single response per document, MLSTM allows multiple outcomes and introduces a hierarchical regression structure to share information across them.
The package provides efficient variational inference algorithms implemented in C++ via Rcpp, enabling scalable estimation for large text corpora.
# install.packages("remotes")
remotes::install_github("thimeno1993/mlstm")
library(mlstm)
set.seed(123)
D <- 50
V <- 200
K <- 5
NZ_per_doc <- 20
NZ <- D * NZ_per_doc
count <- cbind(
d = rep(0:(D - 1), each = NZ_per_doc),
v = sample.int(V, NZ, replace = TRUE) - 1L,
c = rpois(NZ, 3) + 1
)
Y <- cbind(
y1 = rnorm(D),
y2 = rnorm(D)
)
mod_lda <- run_lda_gibbs(
count = count,
K = K,
alpha = 0.1,
beta = 0.01,
n_iter = 20,
verbose = FALSE
)
str(mod_lda$theta)
str(mod_lda$phi)
y <- Y[, 1]
set_threads(2)
mod_stm <- run_stm_vi(
count = count,
y = y,
K = K,
alpha = 0.1,
beta = 0.01,
max_iter = 50,
min_iter = 10,
verbose = FALSE
)
y_hat <- ((mod_stm$nd / mod_stm$ndsum) %*% mod_stm$eta)[, 1]
cor(y, y_hat)
J <- ncol(Y)
mu <- rep(0, K)
upsilon <- K + 2
Omega <- diag(K)
mod_mlstm <- run_mlstm_vi(
count = count,
Y = Y,
K = K,
alpha = 0.1,
beta = 0.01,
mu = mu,
upsilon = upsilon,
Omega = Omega,
max_iter = 50,
min_iter = 10,
verbose = FALSE
)
Y_hat <- ((mod_mlstm$nd / mod_mlstm$ndsum) %*% mod_mlstm$eta)
cor(Y, Y_hat)
Each row of count represents one non-zero document-term entry.
| column | description |
|---|---|
| d | document index (0-based) |
| v | word index (0-based) |
| c | token count |
RcppRcppParallelTomoya Himeno
MIT License
devtools::load_all()
devtools::test()
devtools::check()