SearcharxivSearch

arXiv subjects

Tianxi Cai

Publications and source records attributed to Tianxi Cai.

At least 19 recordsLinked to original sources

Guided Adversarial Robust Transfer Learning with Source Mixing

Transfer learning is a critical technique that enables the application of knowledge gained from existing tasks or domains to improve performance on a new one, reducing the need for extensive data and training in each new context. Many existing transfer learning methods rely on leveraging information from source populations closely resembling the target population. However, this approach often overlooks valuable knowledge that may be present in different yet potentially related auxiliary samples. When dealing with a limited amount of target data and multiple source data, we introduce a novel approach, Guided Adversarial Robust Transfer (GART) learning, that breaks free from strict similarity constraints. GART is designed to optimize the most adversarial loss with respect to a collection of source mixture distributions that guarantee excellent prediction performances for the target data. We establish the closed form of the population GART and show that the GART estimator achieves a faster convergence rate than the model fitted with the target data. Our simulation studies suggest that GART outperforms existing transfer learning methods, attaining higher robustness and accuracy. We highlight GART's predictiveness and robustness by applying it to form genetic prediction models of high-density lipoprotein cholesterol using multi-institutional biobank-linked electronic health records data.

cs.LG

Generalized Linear Markov Decision Process

Offline reinforcement learning for longitudinal studies often faces two linked challenges: rewards may be binary or bounded, and reward observations may be available only for a subset of trajectories or time points even when the corresponding state-action-next-state histories are available. Linear Markov decision process methods are tractable because Bellman backups remain linear, but they require linear rewards and do not indicate how transition-only observations should be used. We introduce GRASP-MDP, Generalized Reward And Semi-supervised Pessimism for Markov Decision Processes, a reward-transition separated framework addressing both issues through one Bellman decomposition. It preserves linear transition dynamics while modeling reward means through generalized linear models. Although this breaks the usual linear Bellman form, the backup remains explicit as a nonlinear reward component plus a linear continuation component, yielding a nonlinear-plus-linear Bellman-complete class. In the resulting recursion, observed rewards estimate the generalized reward component, while all available transitions estimate the continuation component. The resulting pessimistic value iteration controls the two estimation errors separately without reward imputation. Finite-sample guarantees show that transition-only observations reduce transition-estimation error while reward uncertainty remains governed by observed rewards. Simulations and a multiple sclerosis electronic health record application illustrate the empirical benefit of retaining transition-only observations.

stat.ML

Efficient Modeling of Surrogates to Improve Multi-source High-dimensional Integrative Regression

Surrogate variables play an important role in various fields due to the scarcity or absence of gold standard labels. We develop a novel approach named SASH for Surrogate-Assisted and data-Shielding High-dimensional integrative regression. It is a semi-supervised approach that efficiently leverages sizable unlabeled samples with error-prone surrogate outcomes from multiple local sites, to improve the learning accuracy of the small gold-labeled data. To facilitate stable and efficient knowledge extraction from the surrogates, our method first obtains a preliminary supervised estimator, and then uses it to assist training a regularized single index model (SIM) for the surrogates. Interestingly, through a chain of convex and properly penalized sparse regressions that approximate the SIM loss with bias-correction, our method avoids the local minima issue of the SIM, and fully eliminates the impact of the preliminary estimator's excessive error. In addition, it protects individual-level information through the aggregation of summary statistics from local sites, leveraging a similar idea of bias-corrected approximation. Through simulation studies, we demonstrate that our method outperforms existing approaches. Finally, we apply our method to develop a genetic risk model for type II diabetes using large-scale data sets from UK and Mass General Brigham biobanks, where only a small fraction of subjects in one site are labeled through chart reviewing.

stat.ME

Calibrated Estimation and Inference for Semiparametric Regression Models

We consider a broad class of semiparametric regression models in which the conditional distribution of the response takes the form $f\{Y|\boldsymbol{x}^T\boldsymbolβ+m(z),ϕ\}$, known up to a parametric component $\boldsymbolβ$ of diverging dimension $p$, a smooth function $m(\cdot)$, and a dispersion parameter $ϕ$. The existing literature on such models has focused on semiparametric efficiency for $\boldsymbolβ$, treating $ϕ$ and $m(\cdot)$ as nuisances and largely ignoring finite-sample bias. Yet this bias can be substantial, particularly when $p$ is large relative to $n$ or the dispersion is high, and it can seriously undermine inference for $\boldsymbolβ$; moreover, $ϕ$ is often of direct scientific interest. We therefore propose SABRE, a general calibration framework for semiparametric estimation and inference, which calibrates an initial estimator against its model-implied expectation under a tractable parametric approximation to the semiparametric model. For generalized partially linear models, we show that SABRE reduces the bias of both $\boldsymbolβ$ and $ϕ$, accommodates a diverging parameter dimension without sparsity, and preserves the first-order variance and semiparametric efficiency of the initial estimator; the joint construction also improves estimation and inference for $m(\cdot)$. Simulation studies and an application to Alzheimer's disease genetics association analysis demonstrate the empirical effectiveness of SABRE in reducing bias and improving inference.

stat.ME

Multimodal domain adaptation under label shift and blockwise missing modalities

Multimodal domain adaptation uses labeled source datasets to predict outcomes in an unlabeled target population. Here, different sources may observe different subsets of modalities and have distribution shifts from the target. In this scenario, the cross-modal patterns used for alignment depend on the outcome distribution, so aligning unadjusted source data can misrepresent the target. Nevertheless, the distribution shift of the outcome cannot be measured directly when labels are unavailable in the target sample and the missing modality blocks prevent simple pooling. We propose a reference-anchored domain adaptation method. A reference modality observed in every source and the target is used to estimate the target outcome distribution and reweight source observations before alignment. The auxiliary modalities are then mapped to a common, target-defined representation obtained by canonical correlation analysis (CCA) in the target and reproduced in each source by ridge-regression maps. Outcome information is transferred through density-ratio models on the aligned representation. When gold-standard source labels are sparse, a surrogate-label-assisted approach is developed to enable robust domain adaptation. We establish a unified-rotation match-up result for the target CCA and source ridge maps and consistency of the target conditional outcome distribution. Simulations and a renal cell carcinoma (RCC) application show improved calibration and stable prediction under distributional shift and blockwise missing modalities.

stat.ME

Unveiling Invariant and Transferable Latent Factors Across Heterogeneous Environments via ATLAS

This paper considers a multi-environment factor model in which high-dimensional covariates are collected from heterogeneous environments, with auxiliary labels available in a subset of these environments. The joint distribution of the covariates may vary across environments, whereas the latent structure is decomposed into invariant factors with shared loadings and heterogeneous factors with environment-specific loadings. Such a model is motivated by transfer learning and latent factor regression, where one seeks stable low-dimensional representations for both interpretation and robust out-of-sample prediction of the response $Y$. Leveraging the invariance principle, we show that the invariant and heterogeneous factors are disentangled under a minimal structural condition. Based on this, we propose ATLAS, an Auxiliary-label and invariance-guided Transfer via Latent Alignment across heterogeneous environmentS. ATLAS is a unified procedure that leverages the invariance principle to separate aligned invariant and unaligned heterogeneous factors, and further exploits supervision from auxiliary labels to extract prediction-invariant and transferable factors from those unaligned heterogeneous factors. ATLAS yields near-oracle performance for downstream latent factor regression, enables transferable prediction in new environments through the full latent signal when auxiliary labels are available, and reduces to robust invariant-factor-only prediction otherwise. We establish sharp non-asymptotic error bounds for recovering invariant and heterogeneous factors, identifying all the response-invariant factors, and estimating the invariant signal in $Y$.

math.ST

Domain Adaptation Targeting Heterogeneous and Imbalanced Subgroups

Domain adaptation enables generalizable and efficient data-driven research. However, existing work has largely focused on domain adaptation for some intrinsically homogeneous target cohort, overlooking inherent heterogeneity within the target, which can exacerbate biases and unfairness in the presence of subgroups with imbalanced sample sizes. We develop a novel domain adaptation framework that addresses more complicated target data that consists of heterogeneous and data-sparse subgroups and lacks gold-standard label observations. Our method simultaneously handles high-dimensionality, covariate shift, and outcome model heterogeneity by combining a model-assisted debiasing step used for covariate shift correction with an adaptive knowledge-guided sparsification procedure used to mitigate the issue of sample disparity. We also introduce a new model selection strategy to avoid negative knowledge transfer in the absence of labels in the target data. Our method is theoretically justified for being robust to nuisance model misspecification and adaptive to heterogeneity between the subgroups. Numerical experiments and two real-world applications, including genetic risk modeling of type II diabetes and prediction of mutation-induced protein stability changes, demonstrate the practical advantages of our method.

stat.ME

Contrastive Learning on Multimodal Analysis of Electronic Health Records

Electronic health record (EHR) systems capture a wealth of multimodal clinical data, encompassing both structured clinical codes and unstructured clinical notes. Yet, many EHR-focused studies have traditionally examined these modalities in isolation or combined them using simplistic methods, overlooking the intrinsic synergy between them. In reality, these modalities are deeply interconnected, each containing clinically relevant and complementary information that, when integrated effectively, can provide a more comprehensive understanding of patient health. Despite the success of multimodal contrastive learning in vision-language applications, its potential remains under-explored in multimodal EHR, particularly in terms of theoretical understanding. To support statistical analysis of multimodal EHR data, we propose a multimodal feature embedding generative model and design a multimodal contrastive loss to learn EHR feature representations. Our theoretical analysis demonstrates the effectiveness of multimodal learning over single-modality learning and connects the solution of the loss function to the singular value decomposition of a pointwise mutual information matrix. This connection leads to a privacy-preserving algorithm tailored for multimodal EHR representation learning. Simulation studies show that the proposed algorithm performs well under a variety of configurations. We further validate its clinical utility using real-world EHR data.

stat.ML

Spherical Mixture Integration for Latent Embedding Alignment across Multi-Source Feature Spaces

Multi-institutional electronic health record (Multi-EHR) data have emerged as a powerful resource for developing predictive models to support clinical decisions and for generating reliable real-world evidence. By aggregating information from diverse patient populations and institutions, they enhance the robustness and generalizability of models and findings. However, analyzing multi-EHR remains challenging because disparate institutions rarely map all data elements to common ontologies, and raw EHR codes are often overly granular and institution-specific, fragmenting representations of the same clinical concept. Hence, integrative analysis must overcome two key hurdles: harmonizing codes with the same clinical meaning (synonymy), and aligning institutional feature spaces. To address these challenges, we propose SMILE, a Spherical Mixture Integration for Latent Embedding alignment across multi-source feature spaces, where embeddings from heterogeneous sources serve as privacy-preserving summaries of clinical concepts and sparse relational pairs provide weak supervision. Synonymy is modeled via a mixture of von Mises-Fisher distributions, yielding unified representations of semantically equivalent raw codes. We develop a composite quasi-likelihood estimator with non-asymptotic error bounds for the latent representations and mixture mean directions and consistent synonym-cluster recovery, quantifying the gains from integrating multiple sources and knowledge-graph information. Simulations and a multi-institutional EHR application demonstrate improved alignment and synonym clustering.

stat.ME

Enhancing Spectral Embedding through Robust and Flexible Knowledge Transfer in Electronic Health Records

We propose a spectral-based, unsupervised representation learning framework to derive low-dimensional embeddings for clinical concepts and patients in rare disease cohorts from electronic health records, where data are high-dimensional but sample sizes are limited. To overcome this challenge, we incorporate a knowledge matrix extracted from a broader population that shares a partially overlapping subspace with the rare-disease cohort. Our method departs from existing approaches by relaxing restrictive one-to-one signal-alignment assumptions between the latent data matrix and knowledge matrix, allowing more flexible and realistic forms of structured sharing. We introduce a novel two-step spectral embedding procedure: first, we identify and remove irrelevant components from the knowledge matrix; then, we apply a projection-based method to separately recover shared and heterogeneous components. Simulations and an analysis of a real-world multiple sclerosis cohort show that the proposed method outperforms competing approaches, particularly in challenging scenarios where shared signals are weak and only partially aligned, as is common in rare-disease data.

stat.ML

Optimally taming biases in black-box models for efficient semiparametric estimation

Modern semiparametric estimation often relies on flexible black-box machine learning methods to estimate nuisance functions, raising a fundamental question: how do nuisance estimation errors propagate into inference for low-dimensional target parameters? The dominant paradigm, exemplified by double machine learning (DML), yields error bounds in which nuisance estimation errors enter multiplicatively. While widely adopted, it remains unclear whether this multiplicative-rate dependence is optimal for black-box models. In this paper, we start by revisiting the partial linear model $Y = μ_0(X)+T\cdotβ_0+\varepsilon$ under a structure-agnostic setting, where the nuisance function $μ_0$ is estimated using a generic machine learning model, with approximation error $δ^a_μ$ and stochastic error $δ_μ^s$. We show that the standard DML rate is not optimal in the regime where the auxiliary function $\mathbb{E}[T|X=x]$ cannot be consistently estimated. We propose a new estimator for $β_0$ that achieves a sharper rate of $n^{-1/2}+δ^a_μ+(δ_μ^s)^2$ and establish a matching lower bound demonstrating its optimality. Our results reveal a new principle: the first-order stochastic error of nuisance estimation can be eliminated without imposing any additional assumptions. This also leads to a revised tuning strategy favoring under-smoothing, where $δ^a_μ\asymp(δ_μ^s)^2$, rather than the classical bias-variance trade-off $δ^a_μ\asymp δ_μ^s$. Under mild additional conditions, the estimator is asymptotically normal with minimal asymptotic variance. The proposed method extends to a broad class of semi-parametric linear functional estimation problems, including average treatment effect estimation. Our results imply that popular orthogonal score methods in semiparametric estimation with black-box nuisance learners can be substantially improved.

math.ST

Cost-optimal Sequential Testing via Doubly Robust Q-learning

Clinical decision-making often involves selecting tests that are costly, invasive, or time-consuming, motivating individualized, sequential strategies for what to measure and when to stop ascertaining. We study the problem of learning cost-optimal sequential decision policies from retrospective data, where test availability depends on prior results, inducing informative missingness. Under a sequential missing-at-random mechanism, we develop a doubly robust Q-learning framework for estimating optimal policies. The method introduces path-specific inverse probability weights that account for heterogeneous test trajectories and satisfy a normalization property conditional on the observed history. By combining these weights with auxiliary contrast models, we construct orthogonal pseudo-outcomes that enable unbiased policy learning when either the acquisition model or the contrast model is correctly specified. We establish oracle inequalities for the stage-wise contrast estimators, along with convergence rates, regret bounds, and misclassification rates for the learned policy. Simulations demonstrate improved cost-adjusted performance over weighted and complete-case baselines, and an application to a prostate cancer cohort study illustrates how the method reduces testing cost without compromising predictive accuracy.

stat.ML

Representation learning to advance multi-institutional studies with electronic health record data from US and France

The widespread adoption of electronic health records has created new opportunities for translational clinical research, yet this promise remains constrained by fragmented data across privacy-siloed institutions and substantial heterogeneity in local coding practices. While privacy-preserving collaborative learning allows institutions to work together without sharing patient-level data, it does not address inconsistencies in how clinical concepts are represented across sites. We introduce a graph-based framework that addresses this gap by treating data harmonization as a scalable representation learning problem. Rather than relying on fixed standards or manual mappings, the framework integrates institution-specific summary statistics from health records, curated biomedical knowledge graphs, and semantic information derived from large language models to learn a shared semantic space. This joint learning approach aligns diverse, site-specific vocabularies while preserving patient privacy. Evaluated across seven institutions and two languages, the framework provides a robust, data-centric foundation for training and deploying clinical models across heterogeneous healthcare systems.

cs.AI

Nonparametric estimation of the total treatment effect with multiple outcomes in the presence of terminal events

As standards of care advance, patients are living longer and once-fatal diseases are becoming manageable. Clinical trials increasingly focus on reducing disease burden, which can be quantified by the timing and occurrence of multiple non-fatal clinical events. Most existing methods for the analysis of multiple event-time data require stringent modeling assumptions that can be difficult to verify empirically, leading to treatment efficacy estimates that forego interpretability when the underlying assumptions are not met. Moreover, many methods do not appropriately account for informative terminal events, such as premature treatment discontinuation or death, which prevent the occurrence of subsequent events. To address these limitations, we derive and validate estimation and inference procedures for the area under the mean cumulative function (AUMCF), an extension of the restricted mean survival time to the multiple event-time setting. The AUMCF is clinically interpretable, properly accounts for terminal competing risks, and can be estimated nonparametrically. To enable covariate adjustment, we also develop an augmentation estimator that provides efficiency at least equaling, and often exceeding, the unadjusted estimator. The utility and interpretability of the AUMCF are illustrated with extensive simulation studies and through an analysis of multiple heart-failure-related endpoints using data from the Beta-Blocker Evaluation of Survival Trial (BEST) clinical trial. Our open-source R package MCC makes conducting AUMCF analyses straightforward and accessible.

stat.ME

Controllable Sequence Editing for Biological and Clinical Trajectories

Conditional generation models for longitudinal sequences can produce new or modified trajectories given a conditioning input. However, they often lack control over when the condition should take effect (timing) and which variables it should influence (scope). Most methods either operate only on univariate sequences or assume that the condition alters all variables and time steps. In scientific and clinical settings, interventions instead begin at a specific moment, such as the time of drug administration or surgery, and influence only a subset of measurements while the rest of the trajectory remains unchanged. CLEF learns temporal concepts that encode how and when a condition alters future sequence evolution. These concepts allow CLEF to apply targeted edits to the affected time steps and variables while preserving the rest of the sequence. We evaluate CLEF on 8 datasets spanning cellular reprogramming, patient health, and sales, comparing against 9 state-of-the-art baselines. CLEF improves immediate sequence editing accuracy by 16.28% (MAE) on average against their non-CLEF counterparts. Unlike prior models, CLEF enables one-step conditional generation at arbitrary future times, outperforming their non-CLEF counterparts in delayed sequence editing by 26.73% (MAE) on average. We test CLEF under counterfactual inference assumptions and show up to 62.84% (MAE) improvement on zero-shot conditional generation of counterfactual trajectories. In a case study of patients with type 1 diabetes mellitus, CLEF identifies clinical interventions that generate realistic counterfactual trajectories shifted toward healthier outcomes.

cs.LG

Transfer Learning with Network Embeddings under Structured Missingness

Modern data-driven applications increasingly rely on large, heterogeneous datasets collected across multiple sites. Differences in data availability, feature representation, and underlying populations often induce structured missingness, complicating efforts to transfer information from data-rich settings to those with limited data. Many transfer learning methods overlook this structure, limiting their ability to capture meaningful relationships across sites. We propose TransNEST (Transfer learning with Network Embeddings under STructured missingness), a framework that integrates graphical data from source and target sites with prior group structure to construct and refine network embeddings. TransNEST accommodates site-specific features, captures within-group heterogeneity and between-site differences adaptively, and improves embedding estimation under partial feature overlap. We establish the convergence rate for the TransNEST estimator and demonstrate strong finite-sample performance in simulations. We apply TransNEST to a multi-site electronic health record study, transferring feature embeddings from a general hospital system to a pediatric hospital system. Using a hierarchical ontology structure, TransNEST improves pediatric embeddings and supports more accurate pediatric knowledge extraction, achieving the best accuracy for identifying pediatric-specific relational feature pairs compared with benchmark methods.

stat.ME

Knowledge-Embedded Latent Projection for Robust Representation Learning

Latent space models are widely used for analyzing high-dimensional discrete data matrices, such as patient-feature matrices in electronic health records (EHRs), by capturing complex dependence structures through low-dimensional embeddings. However, estimation becomes challenging in the imbalanced regime, where one matrix dimension is much larger than the other. In EHR applications, cohort sizes are often limited by disease prevalence or data availability, whereas the feature space remains extremely large due to the breadth of medical coding system. Motivated by the increasing availability of external semantic embeddings, such as pre-trained embeddings of clinical concepts in EHRs, we propose a knowledge-embedded latent projection model that leverages semantic side information to regularize representation learning. Specifically, we model column embeddings as smooth functions of semantic embeddings via a mapping in a reproducing kernel Hilbert space. We develop a computationally efficient two-step estimation procedure that combines semantically guided subspace construction via kernel principal component analysis with scalable projected gradient descent. We establish estimation error bounds that characterize the trade-off between statistical error and approximation error induced by the kernel projection. Furthermore, we provide local convergence guarantees for our non-convex optimization procedure. Extensive simulation studies and a real-world EHR application demonstrate the effectiveness of the proposed method.

cs.LG

Learning Sequential Decisions from Multiple Sources via Group-Robust Markov Decision Processes

We often collect data from multiple sites (e.g., hospitals) that share common structure but also exhibit heterogeneity. This paper aims to learn robust sequential decision-making policies from such offline, multi-site datasets. To model cross-site uncertainty, we study distributionally robust MDPs with a group-linear structure: all sites share a common feature map, and both the transition kernels and expected reward functions are linear in these shared features. We introduce feature-wise (d-rectangular) uncertainty sets, which preserve tractable robust Bellman recursions while maintaining key cross-site structure. Building on this, we then develop an offline algorithm based on pessimistic value iteration that includes: (i) per-site ridge regression for Bellman targets, (ii) feature-wise worst-case (row-wise minimization) aggregation, and (iii) a data-dependent pessimism penalty computed from the diagonals of the inverse design matrices. We further propose a cluster-level extension that pools similar sites to improve sample efficiency, guided by prior knowledge of site similarity. Under a robust partial coverage assumption, we prove a suboptimality bound for the resulting policy. Overall, our framework addresses multi-site learning with heterogeneous data sources and provides a principled approach to robust planning without relying on strong state-action rectangularity assumptions.

stat.ME