--- title: "Introduction to orthoMTL: Multi-Task Survival Analysis" author: "Kevin Vervier" date: "`r Sys.Date()`" output: rmarkdown::html_vignette vignette: > %\VignetteIndexEntry{Introduction to orthoMTL: Multi-Task Survival Analysis} %\VignetteEngine{knitr::rmarkdown} %\VignetteEncoding{UTF-8} --- ```{r setup, include = FALSE} knitr::opts_chunk$set( collapse = TRUE, comment = "#>", fig.width = 7, fig.height = 5, fig.align = "center" ) library(orthoMTL) ``` ## The Problem In oncology and other therapeutic areas, genomic features associated with patient outcomes often have **time-varying effects**. A mutation that predicts early disease progression may be irrelevant for long-term survival — and vice versa. Standard survival models like the Cox proportional hazards model assume that each feature's effect is constant over time. When this assumption is violated, important signals can be missed. **orthoMTL** addresses this by reframing survival analysis as a **multi-task learning** problem. Instead of fitting a single model for time-to-event, we define binary classification tasks at multiple time thresholds: *"Is this patient progression-free at 4 months? At 6 months? At 10? At 15?"* Each threshold becomes a task, and the model learns a separate coefficient vector for each — while encouraging the coefficient vectors to be structured through an orthogonality penalty. ### The Objective Function The optimisation problem solved by `orthoMTL()` is: $$ \min_W \; \frac{1}{2n}\|XW - Y\|^2_{\text{obs}} \;+\; \lambda\Bigl[\frac{1-\alpha}{2}\,\Omega_K(W)^2 \;+\; \alpha\,\|W\|_1\Bigr], \qquad \Omega_K(W)^2 = \sum_{s,t} K_{st}\,|W_s^\top W_t| $$ where: - $X$ is the $n \times p$ feature matrix - $W$ is the $p \times T$ coefficient matrix ($T$ tasks) - $Y$ is the $n \times T$ label matrix (1 = progression-free, 0 = event, NA = censored) - $\|\cdot\|_{\text{obs}}$ means the loss is only computed on non-censored entries - $K$ is a constraint matrix encoding task relationships — nearby thresholds share more support, distant thresholds are pushed toward orthogonality - $\lambda$ controls the overall penalty strength - $\alpha \in [0, 1]$ is the elastic-net mixing: $\alpha = 0$ gives the pure orthogonality penalty, $\alpha = 1$ gives a pure Lasso ## Simulated Data We begin with simulated data where the **ground truth is known**. The `simulate_mtl()` function generates survival data with five types of time-varying effects: *early*, *late*, *constant*, *increasing*, and *decreasing*. ```{r simulate} set.seed(42) sim <- simulate_mtl( n = 300, p = 15, n_signals = 5, thresholds = c(4, 8, 14, 20), effect_strength = 1.2 ) sim ``` The signal features and their true temporal effect patterns: ```{r ground-truth} # True coefficients for signal features gt <- sim$ground_truth signal_coefs <- gt$coefficients[gt$signal_features, ] signal_coefs ``` ```{r ground-truth-types} # Effect types gt$effect_types ``` ## Data Preparation orthoMTL requires three matrices derived from the survival data. ### Longitudinal binary labels Each patient's survival time and event indicator are converted into a binary label at each threshold. Patients who experienced an event before a threshold are labelled `0`. Patients censored before a threshold have an unknown label (`NA`). All others are labelled `1` (progression-free). ```{r prepare-data} thresholds <- sim$thresholds Y <- create_longitudinal_labels(sim$SurvTime, sim$Event, thresholds) head(Y, 10) ``` ### Censoring indicator matrix A binary matrix marking which labels are observed (`1`) versus censored (`0`). This is used by the solver to mask censored entries in the loss. ```{r indicator} W <- create_indicator_matrix(Y) cat("Proportion observed per threshold:\n") colMeans(W) ``` ### Constraint matrix The diffusion constraint matrix $K$ encodes the prior that nearby thresholds should share coefficient support while distant thresholds can diverge. ```{r constraint} K <- create_constraint_matrix(length(thresholds)) K ``` ## Model Fitting We fit an initial model with a single set of hyperparameters. ```{r fit-initial} fit <- orthoMTL( X = sim$X, Y = Y, lambda = 1e-3, step_size = 0.5, K = K, survival = TRUE, censored.mat = W, alpha = 0.5 ) summary(fit) ``` The coefficient heatmap shows each feature's weight across time thresholds. Blue indicates a protective effect (associated with longer progression-free survival); red indicates a risk effect. ```{r heatmap-initial, fig.width=8, fig.height=6} plot_heatmap(fit) ``` Even with untuned hyperparameters, temporal patterns begin to emerge. ## Cross-Validation We use `cv_orthoMTL()` to search over a grid of hyperparameters and select the configuration with the highest cross-validated C-index. ```{r cv, results='hide'} n_tasks <- length(sim[["thresholds"]]) folds <- rep(1:5, length.out = nrow(sim[["X"]])) cv_res <- cv_orthoMTL( X.train = sim[["X"]], Y.train = Y, W.train = W, K = K, lambdas = c(1e-5, 1e-4, 1e-3), alphas = c(0, 0.5, 1), stepsizes = c(1, 2, 5), diag_vals = c(0.5, n_tasks, 2 * n_tasks), survival = TRUE, folds = folds, n_cores = 1, seed = 42, verbose = FALSE ) ``` ```{r cv-results} print(cv_res) ``` ## Final Model We refit on the full dataset using the best hyperparameters from cross-validation. ```{r fit-final} best <- cv_res$best K_final <- K diag(K_final) <- best$diag_val fit_final <- orthoMTL( X = sim$X, Y = Y, lambda = best$lambda, alpha = best$alpha, step_size = best$stepsize, K = K_final, survival = TRUE, censored.mat = W ) summary(fit_final) ``` ```{r heatmap-final, fig.width=8, fig.height=6} plot_heatmap(fit_final) ``` The task correlation map shows how similar the coefficient profiles are across thresholds. Low distance (red) between adjacent thresholds indicates smooth temporal evolution; high distance (blue) indicates divergent coefficient structures. ```{r correlation, fig.width=6, fig.height=5} plot_correlation(fit_final) ``` ## Bootstrap Inference To assess whether each feature's coefficients are distinguishable from noise, we compare bootstrapped models (resampled data, real signal) against null models (permuted outcomes, no signal). ```{r bootstrap, results='hide'} boot_res <- bootstrap_orthoMTL( X = sim$X, Y = Y, lambda = best$lambda, alpha = best$alpha, step_size = best$stepsize, K = K_final, survival = TRUE, censored.mat = W, n_repeats = 200, n_cores = 1, verbose = FALSE ) ``` ```{r bootstrap-print} print(boot_res) ``` We select four features for detailed inspection: one true signal of each pattern, and one null feature. ```{r bootstrap-select} signal_feats <- gt[["signal_features"]] null_feats <- gt[["null_features"]] effect_types <- gt[["effect_types"]] selected <- c( signal_feats[effect_types == "switch"][1], signal_feats[effect_types == "constant"][1], signal_feats[effect_types == "early"][1], null_feats[1] ) cat("Selected features:\n") cat(" Switch signal: ", selected[1], "(Cox should miss)\n") cat(" Constant signal:", selected[2], "(Cox should find)\n") cat(" Early signal: ", selected[3], "(Cox may dilute)\n") cat(" Null feature: ", selected[4], "(neither should find)\n") ``` ```{r bootstrap-plot, fig.width=8, fig.height=6} suppressWarnings( plot_bootstrap(boot_res, features = selected) ) ``` For the true signal features, the real coefficients (coloured line) separate clearly from the null distribution (grey line). The temporal patterns are visible: the *early* feature's effect is strongest at early thresholds and fades; the *late* feature's effect emerges at later thresholds; the *constant* feature is stable across all thresholds. For the null feature, real and null distributions overlap — the model correctly assigns it no meaningful effect. ## Comparison with Cox Proportional Hazards The Cox model assumes each feature's effect is constant over time. We fit a Cox elastic-net model using `glmnet` and compare which features it detects. ```{r cox-comparison, eval=requireNamespace("glmnet", quietly = TRUE) && requireNamespace("survival", quietly = TRUE)} library(survival) library(glmnet) surv_obj <- Surv(time = sim[["SurvTime"]], event = sim[["Event"]]) # Cross-validate alpha (mixing parameter) alphas <- seq(0, 1, by = 0.1) cv_scores <- numeric(length(alphas)) for (i in seq_along(alphas)) { set.seed(42) cvfit <- cv.glmnet( x = sim[["X"]], y = surv_obj, family = "cox", type.measure = "C", alpha = alphas[i] ) cv_scores[i] <- cvfit[["cvm"]][cvfit[["index"]]["min", ]] } best_alpha <- alphas[which.max(cv_scores)] cat("Best alpha:", best_alpha, "(CV C-index:", max(cv_scores), ")\n") # Refit with best alpha set.seed(42) cox_fit <- cv.glmnet( x = sim[["X"]], y = surv_obj, family = "cox", type.measure = "C", alpha = best_alpha ) cox_coefs <- as.numeric(coef(cox_fit, s = "lambda.min")) names(cox_coefs) <- colnames(sim[["X"]]) ``` ```{r cox-table, eval=requireNamespace("glmnet", quietly = TRUE) && requireNamespace("survival", quietly = TRUE)} # Build comparison table ortho_mean_abs <- apply(abs(coef(fit_final)), 1, mean) comparison <- data.frame( feature = gt$signal_features, effect_type = as.character(gt$effect_types), orthoMTL_mean_abs = round(ortho_mean_abs[gt$signal_features], 4), cox_coef = round(cox_coefs[gt$signal_features], 4), cox_detected = cox_coefs[gt$signal_features] != 0, stringsAsFactors = FALSE ) cat("Signal feature detection comparison:\n\n") print(comparison, row.names = FALSE) ``` The Cox proportional hazards model is a powerful and well-established tool for survival analysis. In terms of overall discrimination (C-index), Cox often matches or exceeds orthoMTL — it directly optimises the survival likelihood, while orthoMTL solves a regression problem on binary labels. The value of orthoMTL is not in replacing Cox but in **complementing** it. Cox produces a single coefficient per feature — an average effect across the entire follow-up. orthoMTL produces a coefficient *per feature per timepoint*, revealing temporal dynamics that a single number cannot capture. ```{r cox-vs-ortho, fig.width=8, fig.height=4, eval=requireNamespace("glmnet", quietly = TRUE) && requireNamespace("survival", quietly = TRUE)} # Pick the switch feature — most interesting temporal pattern switch_feat <- signal_feats[effect_types == "switch"][1] # What Cox sees: one number cat(switch_feat, "— Cox coefficient:", round(cox_coefs[switch_feat], 3), "\n") # What orthoMTL sees: a trajectory cat(switch_feat, "— orthoMTL coefficients:\n") print(round(coef(fit_final)[switch_feat, ], 3)) # What the truth is cat(switch_feat, "— True coefficients (sign-aligned):\n") print(round(-gt[["coefficients"]][switch_feat, ], 3)) ``` Cox reports a single number for this feature. orthoMTL reveals that its effect changes direction over time — information that could guide clinical interpretation of early versus late treatment response. ## Ground Truth Recovery Finally, we compare the estimated coefficients against the true data-generating coefficients for signal features. ```{r ground-truth-recovery} true_coefs <- gt$coefficients[gt$signal_features, ] est_coefs <- coef(fit_final)[gt$signal_features, ] # Sign convention: # Simulation uses log-hazard scale (negative = protective, reduces hazard) # orthoMTL models P(progression-free) (positive = protective) # Negate true coefficients to align true_aligned <- -true_coefs task_cors <- sapply(seq_len(ncol(true_aligned)), function(k) { cor(true_aligned[, k], est_coefs[, k]) }) names(task_cors) <- colnames(true_aligned) cat("Correlation between true and estimated coefficients per threshold:\n") print(round(task_cors, 3)) ``` The simulation generates coefficients on the log-hazard scale (negative = protective), while orthoMTL models the probability of being progression-free (positive = protective). The true coefficients are sign-flipped below so both heatmaps share the same interpretation: blue = protective, red = risk-increasing. ```{r recovery-heatmap, fig.width=8, fig.height=5} make_long <- function(mat, source_label) { data.frame( feature = rep(rownames(mat), ncol(mat)), threshold = rep(colnames(mat), each = nrow(mat)), weight = as.vector(mat), source = source_label, stringsAsFactors = FALSE ) } # Normalize each matrix to [-1, 1] by dividing by its own max absolute value normalize <- function(mat) mat / max(abs(mat), na.rm = TRUE) combined <- rbind( make_long(normalize(true_aligned), "True (sign-aligned)"), make_long(normalize(est_coefs), "Estimated") ) combined$feature <- factor(combined$feature, levels = rev(rownames(true_aligned))) combined$threshold <- factor(combined$threshold, levels = colnames(true_aligned)) combined$source <- factor(combined$source, levels = c("True (sign-aligned)", "Estimated")) # Now both panels use the full color range ggplot2::ggplot(combined, ggplot2::aes(threshold, feature, fill = weight)) + ggplot2::geom_tile() + ggplot2::scale_fill_gradient2(low = "red", mid = "white", high = "blue", midpoint = 0, limits = c(-1, 1)) + ggplot2::facet_wrap(~ source) + ggplot2::labs(x = "Threshold (months)", y = NULL, fill = "Normalized\nCoefficient") + ggplot2::theme_minimal() ``` ## Conclusion This vignette demonstrated the orthoMTL workflow for survival analysis. The key difference from standard Cox modelling is not in overall predictive accuracy but in **interpretability**: orthoMTL reveals how each feature's effect evolves across time thresholds. This is particularly relevant when: - Features have time-varying effects (e.g., early vs late resistance) - The clinical question is not just "does this feature matter?" but "when does it matter?" - Effect reversal over time (switch patterns) may be clinically meaningful but invisible to proportional hazards models For the application of orthoMTL to real clinical data, see the SOLAR-1 analysis in Annals of Oncology (2026). For details on individual functions, see the package help pages (`?orthoMTL`, `?cv_orthoMTL`, `?bootstrap_orthoMTL`, etc.). ### Citation If you use orthoMTL in your work, please cite: > Vervier, K., Mahé, P., d'Aspremont, A., Veyrieras, J.-B., & Vert, J.-P. > (2014). On Learning Matrices with Orthogonal Columns or Disjoint Supports. > *ECML-PKDD 2014*. For the survival extension: > Annals of Oncology (2026). DOI: > [10.1016/j.annonc.2026.04.003](https://doi.org/10.1016/j.annonc.2026.04.003)