Integrated Bayesian machine learning for multi-omics prediction and classification
Wednesday, September 30, 2026
An integrated Bayesian framework for multi-omics prediction and classification (Mallick et al. 2024)
Statistics in Medicine 43(5):983–1002 (2024) doi:10.1002/sim.9953
Goal: find multi-omics biomarkers that predict an outcome (disease status, gestational age, time to labor, …)
Earlier approaches and their gaps:
| Approach | Example | Gap |
|---|---|---|
| Concatenation | random forest on pooled features (Franzosa et al. 2019) | layers differ in size, scale and noise |
| Stacked generalization | elastic net / LASSO stacking (Ghaemi et al. 2019; Stelzer et al. 2021) | no proper CV in the meta-learner → overfitting, leakage |
| sGCCA | DIABLO (Singh et al. 2019) | categorical outcomes only; costly tuning; majority vote |
| Deep learning | “black box”, hard to interpret |
None of these quantify uncertainty of the predictions or feature importance.
Figure: IntegratedLearner GitHub repository (MIT licence). The paper’s default is late fusion with BART base learners.
For \(K\) omics layers measured on the same samples:
The final model is a weighted average of per-layer tree ensembles
A sum of many small trees, each a weak learner (Chipman, George, and McCulloch 2010):
\[ y = \sum_{j=1}^{m} g(\mathbf{x}; T_j, M_j) + \epsilon, \qquad \epsilon \sim N(0, \sigma^2) \]
What BART gives for free:
Stack the \(V\)-fold out-of-sample predictions \(\hat{y}_{ik}\) of each layer \(k\) as “level 1” data.
Continuous outcome: non-negative least squares
\[ \hat{\boldsymbol\alpha} = \arg\min_{\alpha_k \ge 0}\; \sum_{i}\Big(y_i - \sum_{k=1}^{K}\alpha_k\,\hat{y}_{ik}\Big)^2 \]
Binary outcome: minimize the rank loss \(1 - \mathrm{AUC}\) of the combined score \(\sum_k \alpha_k \hat{p}_{ik}\), with \(\alpha_k \ge 0\)
Repeated omics measurements, cross-sectional outcome (e.g. IBD status)
Stage 1: per-feature linear mixed model
\[ x_{ij} = \mathbf{w}_{ij}^T \boldsymbol\beta + \mathbf{z}_{ij}^T \mathbf{b}_i + \epsilon_{ij} \]
Stage 2: use the predicted random effects \(\hat{\mathbf{b}}_i\) as features in IntegratedLearner
Keeps within-subject information that a cross-sectional model would discard.
| Study | Subjects | Samples | Layers | Outcome |
|---|---|---|---|---|
| PRISM (discovery) | 155 | 155 | metabolomics, microbiome | IBD vs. non-IBD |
| PRISM (validation) | 65 | 65 | metabolomics, microbiome | IBD vs. non-IBD |
| iHMP | 132 | 1785 | metabolomics, microbiome | IBD vs. non-IBD (longitudinal features) |
| Pregnancy (Ghaemi et al.) | 17 | 51 | 7 layers | gestational age |
| Labor onset (Stelzer et al.) | 53 + 8 | 150 + 21 | metabolome, proteome, immune | time to labor, gestational age |
Plus simulations from TCGA ovarian cancer multi-omics (InterSIM).
Findings:
TCGA-like multi-omics (InterSIM): 131 genes, 367 methylation sites, 160 proteins; \(n\) = 200, 500, 1000; SNR = 1, 5, 10; 100 replicates
| Scenario | IntegratedLearner vs. stacked linear models |
|---|---|
| Highly nonlinear (Friedman function) | best in all 9 settings (e.g. \(R^2\) 0.81 vs. 0.63 at \(n\) = 1000, SNR 10) |
| Highly linear | best in 5 of 9 settings |
| Fully linear | second best: regularized linear models win; always beats concatenated RF |
Runtime is on par with the frequentist methods, even though it combines MCMC with cross-validation, thanks to the fast bartMachine implementation of BART.
The GitHub version extends the paper’s method:
run_concat), late (run_stacked), and intermediate / cooperative learning (run_intermediate, via multiview)sl_bart (paper default, needs Java ≥ 21), any SuperLearner SL.* learner, and native multiclass and survival learnersTrain on PRISM, validate on an independent cohort (NLIBD)
ExperimentList class object of length 2:
[1] species: SummarizedExperiment with 340 rows and 155 columns
[2] metabolites: SummarizedExperiment with 1500 rows and 155 columns
IBD
0 1
34 121
fit_bin <- IntegratedLearner(
MAE_train = PRISM_MAE,
MAE_valid = NLIBD_MAE,
folds = 5,
base_learner = "sl_bart", # BART base learners (paper default)
meta_learner = "sl_nnls_auc", # non-negative weights, maximize AUC
filter_method = "prevalence", filter_pct = 40,
run_screening = TRUE, screen_pct = 30,
print_learner = FALSE,
family = binomial()
)Also fits the concatenated model (early fusion) for comparison. Filtering and screening are fold-safe (learned only on training folds).
Lesson: always check external validation, and compare the fusion strategies on your own data; IntegratedLearner makes that comparison easy.
7 omics layers, 17 women, repeated samples; continuous outcome
data("pregnancy_MAE", package = "IntegratedLearner")
fit_cont <- IntegratedLearner(
MAE_train = pregnancy_MAE,
folds = 5,
base_learner = "sl_bart",
meta_learner = "sl_nnls_auc",
filter_method = "variance", filter_pct = 40,
run_screening = TRUE, screen_pct = 30,
print_learner = FALSE,
family = gaussian()
) R2 weight
CellfreeRNA 0.14 0.00
ImmuneSystem 0.01 0.00
Metabolomics 0.70 0.44
Microbiome 0.59 0.30
PlasmaLuminex 0.15 0.00
PlasmaSomalogic 0.65 0.27
SerumLuminex 0.00 0.00
stacked 0.80 NA
concatenated 0.42 NA
As in the paper, the stacked model beats every single layer, and metabolomics, microbiome and plasma SomaLogic get the largest weights.
Weight the posterior draws of each layer’s BART model by the layer weights:
post <- lapply(seq_along(fit_cont$weights), function(k) {
m <- fit_cont$model_fits$model_layers[[k]]
X <- fit_cont$X_train_layers[[k]][, m$training_data_features, drop = FALSE]
bartMachine::bart_machine_get_posterior(m, X)$y_hat_posterior_samples
})
post_fused <- Reduce(`+`, Map(`*`, post, fit_cont$weights)) # samples x draws
dim(post_fused)[1] 51 1000
library(bayesplot)
rownames(post_fused) <- rownames(fit_cont$X_train_layers[[1]])
y <- setNames(fit_cont$Y_train, rownames(post_fused))
ord <- names(sort(rowMeans(post_fused)))
mcmc_intervals(t(post_fused), prob = 0.68, prob_outer = 0.95) +
scale_y_discrete(limits = ord) +
geom_point(aes(x = y[ord], y = ord), shape = 1, size = 2) +
coord_flip() +
labs(x = "Gestational age", y = "Samples (fitted, training data)") +
theme_bw(base_size = 14) +
theme(axis.text.x = element_blank())Thick (thin) bars: 68% (95%) credible intervals of the fused prediction; circles: observed values.
top_layer <- names(which.max(fit_cont$weights))
invisible(capture.output( # silence progress output
vi <- bartMachine::investigate_var_importance(
fit_cont$model_fits$model_layers[[top_layer]], plot = FALSE
)
))
top <- head(sort(vi$avg_var_props, decreasing = TRUE), 10)
df <- data.frame(feature = substr(names(top), 1, 45), prop = top,
sd = vi$sd_var_props[names(top)])
ggplot(df, aes(reorder(feature, prop), prop)) +
geom_col(fill = "lightsalmon") +
geom_errorbar(aes(ymin = pmax(prop - sd, 0), ymax = prop + sd), width = 0.2) +
coord_flip() +
labs(x = NULL, y = "Inclusion proportion", title = top_layer) +
theme_bw(base_size = 14)# Swap learners: any SuperLearner model as base / meta learner
IntegratedLearner(MAE_train = mae, base_learner = "SL.randomForest",
meta_learner = "sl_nnls_auc", family = binomial())
# Multiclass outcome (native multiclass backend)
IntegratedLearner(MAE_train = mae, outcome_col = "diseaseCat",
base_learner = "randomforest", meta_learner = "randomforest",
family = binomial())
# Survival outcome
IntegratedLearner(MAE_train = mae, base_learner = "surv.coxph")
# Intermediate fusion: cooperative learning (multiview)
IntegratedLearner(MAE_train = mae, run_intermediate = TRUE,
family = binomial())
# New data with missing layers: re-learn only the layer weights
update.learner(fit, feature_table_valid = ft, feature_metadata_valid = fm,
sample_metadata_valid = sm)Software: https://github.com/himelmallick/IntegratedLearner. Workflow figure from the IntegratedLearner repository (MIT licence).