Causal Longitudinal Prior-Fitted Networks for Counterfactual Outcome Prediction

Amirhossein Zare1 Amirhessam Zare1 Herlock Rahimi2
Reza Salarikia3 Mohammad Kashkooli4,5
amhosseinzare@gmail.com amir.hessam.zare@gmail.com herlock.rahimi@yale.edu
salarikiareza@gmail.com mohammadkashkooli594@gmail.com


Abstract

Longitudinal treatment decisions require predicting potential outcomes under future treatment sequences in the presence of time-varying confounding, heterogeneous patient dynamics, and limited domain-specific data. Existing longitudinal causal estimators typically address this problem by training a new model for each cohort or simulator. We introduce Causal Longitudinal Prior-Fitted Networks (CausalLongPFN), a prior-fitted in-context predictor for longitudinal causal prediction. To our knowledge, CausalLongPFN is the first PFN-style model for history-conditional potential-outcome prediction under planned longitudinal treatment sequences, with systematic comparison against established longitudinal causal baselines on branchable counterfactual treatment-response benchmarks and factual real-world clinical data. The model is pretrained entirely on synthetic episodes sampled from a broad prior over temporal structural causal models, exposing it to treatment–confounder feedback, latent heterogeneity, nonlinear state evolution, delayed effects, and cumulative treatment responses. At test time, CausalLongPFN is frozen: it conditions on support trajectories, a query history, and a proposed future treatment sequence, and returns a predictive distribution over future outcomes without gradient updates or propensity-model fitting. Multi-step predictions are obtained by recursively applying the one-step predictor under the specified treatment sequence. We evaluate on branchable cancer, HIV, and warfarin benchmarks with ground-truth counterfactual labels, and on factual-only rolling-origin prediction in MIMIC-III ICU trajectories. CausalLongPFN is competitive with domain-trained longitudinal baselines on counterfactual benchmarks and performs strongly on factual MIMIC-III prediction, suggesting that broad synthetic causal pretraining can provide a useful frozen alternative when repeated domain-specific training is costly or impractical.

1 Introduction↩︎

Predicting how a longitudinal outcome would evolve under future treatment decisions is a central problem in causal inference from longitudinal observational records. In the potential-outcomes framework [1], for a unit observed through history \(H_t\), the target is a history-conditional potential outcome, such as \(\mathbb{E}[Y_{t+\tau}^{(\bar a)}\mid H_t]\), under a planned treatment sequence \(\bar a=(a_t,\ldots,a_{t+\tau-1})\). Under consistency, positivity, and sequential exchangeability, this quantity is identified by the longitudinal \(g\)-formula [2][4]. In practice, however, estimating it is difficult: treatment assignment at each step depends on covariates that are themselves outcomes of prior treatment (time-varying confounding); errors accumulate over the multi-step rollout; and observational cohorts are often too small to fit reliable deep sequence models from scratch.

Modern longitudinal causal estimators address these challenges by explicitly modeling treatment–confounder feedback. RMSN [5] combines recurrent outcome models with inverse-probability weighting; CRN [6] learns balanced recurrent representations using adversarial treatment prediction; and G-Net [7] implements neural \(g\)-computation through autoregressive simulation. More recent transformer-based methods, including the Causal Transformer [8] and G-Transformer [9], use attention to represent longitudinal histories and have achieved strong performance on standard counterfactual benchmarks. Despite this progress, these methods share a fundamental operational constraint: each new cohort or simulator typically requires a separate supervised training run, including validation-based hyperparameter selection and, for some methods, propensity modeling or representation balancing. This pipeline must be repeated for every new cohort or data release.

Prior-Fitted Networks (PFNs) offer a complementary route. Rather than fitting a new model for each dataset, a PFN is pretrained on tasks sampled from a prior over data-generating processes and then performs in-context prediction on a new dataset without gradient updates [10]. This idea has led to strong amortized predictors for tabular data [11] and time-series forecasting [12], [13]. Recent work has also begun to apply PFN-style models for cross-sectional causal inference, including Do-PFN, CausalPFN, and CausalFM [14][16]. However, existing causal PFNs operate on independent, cross-sectional, tabular data: none model the sequential structure of longitudinal histories, handle time-varying confounding, or support multi-step potential outcomes under future treatment sequences. CausalTimePrior [17] introduced a synthetic prior over temporal SCMs with paired observational and interventional time series and demonstrated a PFN-based proof of concept on held-out temporal SCMs. It is primarily positioned as a generic interventional time-series prior, rather than as an end-to-end PFN-style model for history-conditional potential-outcome prediction under planned longitudinal treatment sequences with systematic comparison against established longitudinal causal baselines on synthetic and real-world treatment-response benchmarks. The intersection of longitudinal causal inference and PFN-style in-context prediction therefore remains largely unexplored.

1.0.0.1 This work.

We introduce CausalLongPFN, a prior-fitted network for multi-step counterfactual outcome prediction from longitudinal observational data. Given support trajectories from a new domain, a query history observed up to time \(t\), and a supplied future treatment sequence, the frozen model returns a predictive distribution for the query outcome under that sequence. It does so without target-domain gradient updates, propensity-model fitting, adversarial balancing, or domain-specific simulator access at test time.

The key idea is to amortize longitudinal causal prediction across a broad prior over temporal structural causal models. During pretraining, each synthetic task contains treatment-confounder feedback, latent unit heterogeneity, nonlinear state dynamics, delayed and cumulative treatment effects, and stochastic outcome mechanisms. The model learns to use support trajectories as an in-context description of the task and to answer query-level potential-outcome questions under proposed treatment sequences. At evaluation time, the learned one-step predictor is composed autoregressively under the supplied treatment sequence, yielding multi-step potential-outcome predictions without retraining.

This framing does not remove the standard assumptions needed to interpret observational data causally. Rather, CausalLongPFN provides an amortized estimator for history-conditional potential outcomes in settings where the relevant longitudinal causal structure is supported by the synthetic prior and the usual identification assumptions are plausible. Empirically, we compare the frozen model against MSM, RMSN, G-Net, CRN, Causal Transformer, and G-Transformer, each trained and tuned separately on the target domain. CausalLongPFN achieves competitive normalized RMSE on branchable counterfactual benchmarks and strong factual prediction on MIMIC-III, suggesting that broad synthetic causal pretraining can be a useful alternative to repeated domain-specific training.

1.0.0.2 Contributions.

  1. A prior-fitted model for longitudinal causal prediction. We propose CausalLongPFN for history-conditional potential-outcome prediction under planned longitudinal treatment sequences. Unlike standard longitudinal causal estimators, it is evaluated as a frozen model and requires no test-time adaptation.

  2. A synthetic prior over longitudinal causal tasks. We design a temporal structural causal model prior that generates diverse longitudinal tasks with treatment-confounder feedback, latent unit heterogeneity, nonlinear lagged dynamics, delayed and cumulative treatment effects, regime changes, and mixed noise mechanisms. This prior supplies the support trajectories and query counterfactual targets used for pretraining.

  3. Architecture for longitudinal in-context causal inference. We propose a dual-encoder architecture combining a causal Transformer history encoder with a PFN context encoder over support trajectories and a Gaussian-mixture prediction head for distributional outcomes.

  4. Autoregressive counterfactual rollout. We extend the learned one-step predictor to multi-step prediction by autoregressively rolling it forward under supplied treatment sequences, using each predicted intermediate outcome as part of the subsequent query history.

  5. Zero-shot evaluation against domain-trained baselines. We evaluate a single frozen CausalLongPFN on branchable cancer, HIV, and warfarin counterfactual benchmarks and on factual MIMIC-III ICU prediction. The comparison contrasts amortized synthetic pretraining with baselines that receive domain-specific training and validation-based model selection.

2 Methods↩︎

2.1 Problem formulation↩︎

We consider longitudinal observational data consisting of repeated covariates, treatments, outcomes, and static features. For unit \(i\) at discrete time \(t\), let \[S_{i,t}\in\mathbb{R}^{d_S},\qquad A_{i,t}\in\mathcal{A},\qquad Y_{i,t}\in\mathbb{R},\qquad C_i\in\mathbb{R}^{d_C}\] denote time-varying covariates, treatment, scalar outcome, and static covariates. The model-facing longitudinal state is \[X_{i,t}=(S_{i,t},Y_{i,t})\in\mathbb{R}^{d},\qquad d=d_S+1.\] Implementation-specific details such as padding, the discrete four-action interface, and inactive dimensions are described in Appendix 9.

We use the standard longitudinal ordering in which covariates and the outcome at time \(t\) are observed before treatment \(A_{i,t}\) is assigned. The observed history available at decision time \(t\) is therefore \[H_{i,t} = \bigl(C_i,X_{i,0},A_{i,0},X_{i,1},\ldots,A_{i,t-1},X_{i,t}\bigr). \label{eq:history}\tag{1}\] A one-step potential outcome from time \(t\) to \(t+1\) is indexed by the candidate treatment \(A_{i,t}\) applied after observing \(H_{i,t}\). Given a last observed time \(t_{\mathrm{obs}}\), a horizon \(\tau\ge 1\), and \(t_\star=t_{\mathrm{obs}}+\tau\), we write the planned future treatment sequence as \[\bar a_{t_{\mathrm{obs}}:t_\star-1} = (a_{t_{\mathrm{obs}}},a_{t_{\mathrm{obs}}+1},\ldots,a_{t_\star-1}).\] The first planned treatment is applied after observing \(X_{t_{\mathrm{obs}}}\).

The prediction target is the conditional counterfactual predictive distribution for a query unit, \[p\!\left( Y^q_{t_\star}(\bar a_{t_{\mathrm{obs}}:t_\star-1})\in dy \mid H^q_{t_{\mathrm{obs}}},\mathcal{C} \right), \label{eq:target95predictive}\tag{2}\] where \(\mathcal{C}\) denotes support trajectories from the same task or domain. The support trajectories provide task-specific information about the longitudinal data-generating process, while the query history specifies the individual whose future outcome is to be predicted. At prediction time, the model observes the query history through \(t_{\mathrm{obs}}\) and the planned future treatments. Future query covariates are not observed under the hypothetical treatment sequence and are therefore excluded from the query information set. For multi-step prediction, future query outcomes are generated recursively by the model itself.

For real observational data, a causal interpretation of Eq. 2 requires the usual longitudinal assumptions: consistency, positivity, and sequential exchangeability conditional on the measured history [2][4]. Under these assumptions, the corresponding counterfactual mean is identified by the longitudinal \(g\)-formula. CausalLongPFN does not fit a separate propensity model, balancing representation, or outcome model for each target domain. Instead, it amortizes the prediction of history-conditional potential outcomes by training a prior-fitted network on synthetic longitudinal causal tasks sampled from a broad prior over temporal structural causal models (TSCMs), following the structural-causal-model perspective on interventions and counterfactuals [18], [19]. As in prior-fitted networks, task adaptation occurs through conditioning on the support trajectories in context, rather than through test-time gradient updates [10], [11], [20].

2.2 Causal Longitudinal PFN↩︎

2.2.0.1 Overview.

CausalLongPFN combines a synthetic prior over longitudinal causal tasks with an in-context transformer predictor. During pretraining, each task is sampled from a TSCM prior and provides support trajectories together with query-level factual or counterfactual prediction targets. After pretraining, the model is kept frozen. At test time, it receives support trajectories from a new domain, a query history through \(t_{\mathrm{obs}}\), and a planned future treatment sequence, and returns a predictive distribution for the query outcome under that sequence. Thus, the model is designed as an amortized estimator for longitudinal potential-outcome prediction rather than as a domain-specific supervised learner.

2.2.0.2 Temporal structural causal prior.

Each training episode samples a temporal structural causal model (TSCM) \(\mathcal{M}\sim\Pi\) and then draws support and query trajectories from this sampled data-generating process. The prior is designed to span a broad class of longitudinal causal dynamics rather than to reproduce a single hand-built simulator. A sampled TSCM specifies:

  1. Causal temporal graph. The latent longitudinal state \(S_t\in\mathbb{R}^{d_S}\) has variable dimension and evolves according to sparse contemporaneous and lagged dependencies. Within a time slice, the instantaneous graph is acyclic; across time, lagged edges induce temporal dependence across time. This exposes the model to settings in which current covariates depend on previous covariates, previous treatments, and other variables in the same time slice.

  2. Nonlinear structural mechanisms. State coordinates follow sparse nonlinear autoregressive updates with randomly sampled elementary nonlinearities, including identity, \(\tanh\), sinusoidal, rectified, absolute-value, square, and softplus functions, and with Gaussian, uniform, Laplace, or zero noise. The prior therefore includes both smooth and nonsmooth dynamics, low- and moderate-noise regimes, and occasional nonstationarity through regime switches. Full sampling details are given in Appendix 7.

  3. Longitudinal dynamical motifs. In addition to generic nonlinear dynamics, the prior optionally overlays structured dynamical motifs on randomly selected state dimensions. These include action-memory, saturating, homeostatic, feedback-control, and smoothed-readout channels. The motifs are intended to capture qualitative mechanisms common in longitudinal data, such as delayed treatment effects, bounded accumulation, regulatory feedback, proxy measurements, and slow physiological responses. Motif equations and parameter ranges are listed in Appendix 7.4.

  4. Confounded behavior policy. Treatments in support trajectories and factual query prefixes are sampled from a state-dependent stochastic behavior policy. Each unit has latent heterogeneity \(Z_i\), which affects both its initial state and its treatment policy. This produces time-varying treatment–confounder feedback with varying strength.

  5. Autoregressive outcome model. The scalar outcome is generated as an autoregressive readout of the evolving state with direct and cumulative treatment effects. Consequently, the target may depend on current state, previous outcomes, treatment history, and accumulated exposure. A regime switch is included in a minority of sampled TSCMs to expose the model to nonstationarity.

For interventional query episodes, the generator first simulates the query factual prefix up to \(t_{\mathrm{obs}}\). It then fixes the future treatment sequence \(\bar a_{t_{\mathrm{obs}}:t_\star-1}\) and replays the structural equations forward from the same query state under this intervention, with future additive noise set to its conditional mean. This produces a structural target for the intervention-specific conditional mean. In observational-mode episodes, the query continues under the behavior policy and the target is factual. Details are given in Appendix 7.7.

2.2.0.3 Support-query episode construction.

A pretraining episode is a supervised in-context prediction problem generated from one sampled TSCM. The episode contains support trajectories, a query trajectory prefix, a planned future treatment sequence for the query, and a target outcome. The support trajectories serve as examples from the same task-specific longitudinal system; the query asks for the outcome of one unit under a factual or hypothetical continuation.

Training uses one-step prediction problems sampled at different depths along a future path. After choosing an observation time \(t_{\mathrm{obs}}\) and a future target window, the generator samples a current rollout time \(r\ge t_{\mathrm{obs}}\) and trains the model to predict \(Y^q_{r+1}\) from the query history through \(r\), the current treatment \(A^q_r\), and the support trajectories. Interventional episodes replace future behavior-policy treatments with a sampled hypothetical sequence, whereas observational episodes retain the factual behavior-policy continuation. Multi-step prediction is therefore not trained as a separate direct-horizon task; it is obtained at test time by recursively applying the learned one-step predictor.

To make the support trajectories informative about the sampled task, each support unit contributes several labeled time points from its observed trajectory. These labels provide in-context examples of how histories and treatments map to subsequent outcomes within the same TSCM. Query variables beyond the information available at the current prediction time are hidden according to the information set defined in Section 2.1. Additional details on support-anchor selection, masking, task-local normalization, and training augmentations are given in Appendix 8.

2.2.0.4 Architecture.

CausalLongPFN has three main components: a causal history encoder, a PFN context encoder, and a distributional prediction head. Architectural details are provided in Appendix 9.

(i) Causal history encoder. A trajectory-level causal transformer, implemented using masked self-attention in the Transformer architecture [21], maps each longitudinal sequence to history representations. The encoder processes covariates, outcomes, treatments, and missingness indicators while using a causal attention mask, so the representation at time \(r\) depends only on information available up to that time.

(ii) PFN context encoder. The PFN context encoder performs in-context adaptation from the support trajectories. Support tokens summarize labeled support histories, while the query token summarizes the query history and planned current treatment. The support and query tokens are processed jointly by self-attention. No positional encoding is assigned to the ordering of support trajectories, so the architecture is designed to treat the support trajectories as an unordered set.

(iii) Gaussian-mixture prediction head. The final query representation parameterizes a Gaussian mixture distribution [22] for the normalized next outcome, \[q_\theta(y_{r+1}\mid\mathcal{C},H^q_r,A^q_r) = \sum_{k=1}^{5} \pi_{r,k}\mathcal{N}(y_{r+1};\mu_{r,k},\sigma^2_{r,k}). \label{eq:gmm95density}\tag{3}\] The mixture head provides both a point prediction, given by the mixture mean, and a predictive distribution for uncertainty evaluation. In the implementation, the component means are residualized around the most recent visible or self-predicted outcome, which gives the model a stable persistence baseline at initialization.

2.2.0.5 Implementation scope.

The implemented CausalLongPFN uses a fixed interface across all tasks. Histories contain up to \(60\) observed time points, rollouts are evaluated up to horizon \(5\), and inputs support up to \(10\) time-varying covariate channels, one scalar outcome channel, \(5\) static covariates, and four discrete treatment actions. The model uses a \(4\)-layer causal history encoder and a \(6\)-layer PFN context encoder with hidden dimension \(256\), \(8\) attention heads, feed-forward width \(1024\), and a \(5\)-component Gaussian-mixture prediction head, giving \(8{,}117{,}519\) trainable parameters. During synthetic pretraining, each episode is generated from a sampled TSCM and contains between \(3\) and \(500\) support trajectories. With \(10{,}000\) optimizer updates and effective batch size \(256\), pretraining processes \(2{,}560{,}000\) independently sampled synthetic episodes. At evaluation time, this same model is frozen and applied to all benchmark domains without architectural changes or gradient updates. Padding, support-anchor construction, normalization, and optimization details are provided in Appendices 9, 8, and 11.

2.2.0.6 Training and rollout.

The model is pretrained on synthetic support-query episodes using a one-step Gaussian-mixture negative log-likelihood. The loss is augmented with a small auxiliary term on the mixture mean and a mild regularizer against premature mixture collapse. Optimization details, including AdamW, learning-rate schedule, gradient accumulation, mixed precision, gradient clipping, and stochastic PFN depth, are reported in Appendices 10 and 11.

At test time, all model parameters are frozen. For one-step prediction, the model directly evaluates \(q_\theta(y_{t_{\mathrm{obs}}+1}\mid\mathcal{C},H^q_{t_{\mathrm{obs}}}, a_{t_{\mathrm{obs}}})\). For a horizon \(\tau>1\), CausalLongPFN performs a plug-in sequential rollout under the supplied treatment sequence. Starting at \(r=t_{\mathrm{obs}}\), it predicts the next-outcome distribution under the planned treatment \(a_r\), inserts the mixture mean as the next query outcome, keeps future query covariates unavailable, and repeats this procedure until \(r=t_\star-1\). The final mixture is reported as the predictive distribution for \(Y_{t_\star}^q(\bar a_{t_{\mathrm{obs}}:t_\star-1})\), and its mean is used for point-estimation metrics.

This rollout is a deterministic plug-in approximation to the full posterior predictive distribution over future outcome paths. It is closely related to sequential g-computation and parametric implementations of the longitudinal g-formula [2], [4]. The learned one-step conditional predictor is composed forward under the specified treatment sequence, with predicted intermediate outcomes becoming part of the subsequent query history. A stochastic ancestral rollout that samples intermediate outcomes from the mixture is a natural extension. Appendix 12 gives the algorithmic details and discusses the deterministic plug-in nature of this approximation.

3 Experiments↩︎

We evaluate whether a single frozen CausalLongPFN pretrained only on synthetic TSCM episodes can serve as an in-context predictor on external longitudinal treatment-response tasks. The central comparison is between amortized synthetic pretraining and domain-specific supervised training: CausalLongPFN is evaluated without updating its parameters, whereas all baselines are trained and selected separately using the support trajectories of each target domain.

3.0.0.1 Benchmarks.

We use four longitudinal benchmarks: cancer tumor growth [5][8], [23], warfarin PK/PD [24][26], HIV treatment dynamics [27], [28], and MIMIC-III ICU trajectories [29][33]. These benchmarks are summarized in Table 8 and described in Appendix 13. Cancer, warfarin, and HIV are branchable simulated or semi-mechanistic systems. For these domains, the same patient-specific dynamics can be replayed under alternative future treatment sequences, giving ground-truth counterfactual outcomes for evaluation. MIMIC-III is a real observational ICU dataset and is therefore used only for factual rolling-origin prediction under the observed future treatments. Its role is to test factual temporal prediction on real clinical trajectories, not to validate individual counterfactual effects under unobserved interventions.

3.0.0.2 Evaluation configuration.

All benchmark domains are mapped to a common longitudinal prediction format. Each trajectory contains up to \(60\) time points, prediction origins are selected only after at least \(10\) observed time points, and multi-step evaluation uses a five-step horizon. For each domain, the evaluation grid crosses five support sizes, \(n_{\mathrm{sup}}\in\{40,80,160,320,500\}\), ten task-index levels, and two random repetitions, yielding \(5\times10\times2=100\) benchmark tasks per domain. Each benchmark task contains multiple rolling-origin query rows, which are first aggregated before domain-level summaries are computed. In cancer, HIV, and warfarin, the ten task-index levels correspond to confounding levels that control the strength of state-dependent treatment assignment. In MIMIC-III, the same ten-level grid is retained only to match the benchmark organization across domains; it indexes factual rolling-origin task variants and does not alter the observed ICU trajectories. Cancer, HIV, and warfarin provide branchable counterfactual labels, whereas MIMIC-III provides factual labels under observed future treatments. Dataset construction details are given in Appendix 13, and scoring details are given in Appendix 15.

3.0.0.3 Baselines.

We compare against six standard longitudinal causal baselines: a marginal structural model (MSM) [3], Recurrent Marginal Structural Networks (RMSN) [5], G-Net [7], Counterfactual Recurrent Networks (CRN) [6], Causal Transformer (CT) [8], and G-Transformer (GT) [9]. Together, these methods represent inverse-probability weighting, recurrent neural \(g\)-computation, adversarial representation balancing, and transformer-based longitudinal counterfactual modeling. Each baseline uses support-set validation for model selection and is then refit on the target support trajectories before query evaluation. In contrast, CausalLongPFN receives the same support trajectories only as in-context input and remains frozen.

3.0.0.4 Prediction protocol.

All methods follow the same observation-time convention from Section 2.1. The query history is observed through \(t_{\mathrm{obs}}\), the first planned treatment is \(a_{t_{\mathrm{obs}}}\), and the target is \(Y_{t_{\mathrm{obs}}+\tau}\). For multi-step prediction, methods are evaluated under the supplied future treatment sequence. In branchable simulated domains, this sequence defines the intervention used to generate the counterfactual label. In MIMIC-III, the sequence is the observed future treatment path and the label is factual. Implementation details of scoring are given in Appendix 15.

3.0.0.5 Metrics.

The primary metric is normalized RMSE. Normalization statistics are computed from support trajectories only, so query targets are never used to define the reporting scale. Metrics are first computed for each benchmark task by aggregating all scored query rows within that task, and are then averaged within domains. Domain-balanced summaries average the four domain means equally, preventing large domains from dominating the reported overall performance. For MIMIC-III, normalized RMSE measures factual temporal prediction under observed clinical practice rather than counterfactual accuracy. Lower values are better; in result tables, the best, second-best, and third-best values within each comparison are highlighted in green, blue, and orange, respectively. Full scoring details are provided in Appendix 15.

3.1 Results↩︎

3.1.0.1 Domain-balanced performance.

Table 1 reports the mean normalized RMSE after first aggregating within each domain and then averaging equally across the four domains. CausalLongPFN achieves the best domain-balanced one-step performance, with normalized RMSE \(0.2217\), narrowly ahead of MSM (\(0.2233\)) and RMSN (\(0.2247\)). For five-step prediction, CausalLongPFN ranks third overall, behind RMSN and G-Net, while remaining ahead of MSM, CRN, GT, and CT. These results show that the frozen synthetically pretrained model is competitive with baselines that are trained and selected separately for each target domain.

7pt

Table 1: Domain-balanced normalized RMSE across cancer, HIV, MIMIC-III, andwarfarin. Entries report mean \(\pm\) standard deviation across domains. Lower isbetter. CausalLongPFN is best for one-step prediction and third forfive-step rollout despite using no domain-specific training.
Method One-step Horizon-5
CausalLongPFN 0.222\(\pm\)​0.269 0.389\(\pm\)​0.214
MSM 0.223\(\pm\)​0.275 0.418\(\pm\)​0.292
RMSN 0.225\(\pm\)​0.273 0.350\(\pm\)​0.254
G-Net 0.247\(\pm\)​0.251 0.379\(\pm\)​0.223
CT 0.258\(\pm\)​0.259 0.871\(\pm\)​0.096
GT 0.272\(\pm\)​0.238 0.489\(\pm\)​0.164
CRN 0.347\(\pm\)​0.184 0.472\(\pm\)​0.188

3.1.0.2 Per-domain results.

Table 2 reports normalized RMSE by domain, prediction task, and method. The main pattern is that CausalLongPFN remains consistently competitive across heterogeneous domains without target-domain retraining. For one-step prediction, it ranks second on cancer, third on HIV, first on MIMIC-III, and second on warfarin. For five-step prediction, it ranks first on MIMIC-III and second on warfarin, but is weaker on HIV and cancer, where domain-trained recurrent baselines perform best. This domain-level breakdown is important: CausalLongPFN is not uniformly superior, but it provides a strong frozen predictor across tasks with very different dynamics and outcome scales.

2.8pt

Table 2: Per-domain normalized RMSE with standard deviation across units.Entries report mean \(\pm\) standard deviation. Lower mean normalized RMSE isbetter. The top three values in each row are highlighted. MIMIC-III isfactual-only; cancer, HIV, and warfarin provide branchable counterfactuallabels.
Domain Task CausalLongPFN MSM RMSN G-Net CRN CT GT
Cancer One-step 0.167\(\pm\)​0.255 0.200\(\pm\)​0.278 0.166\(\pm\)​0.256 0.168\(\pm\)​0.242 0.251\(\pm\)​0.291 0.209\(\pm\)​0.265 0.217\(\pm\)​0.281
Cancer Horizon-5 0.385\(\pm\)​0.356 0.465\(\pm\)​0.435 0.246\(\pm\)​0.337 0.308\(\pm\)​0.285 0.278\(\pm\)​0.334 0.849\(\pm\)​0.743 0.372\(\pm\)​0.456
HIV One-step 0.066\(\pm\)​0.032 0.061\(\pm\)​0.027 0.051\(\pm\)​0.029 0.097\(\pm\)​0.066 0.244\(\pm\)​0.175 0.100\(\pm\)​0.058 0.094\(\pm\)​0.056
HIV Horizon-5 0.288\(\pm\)​0.174 0.248\(\pm\)​0.122 0.186\(\pm\)​0.117 0.235\(\pm\)​0.137 0.405\(\pm\)​0.253 0.915\(\pm\)​0.618 0.342\(\pm\)​0.193
MIMIC-III One-step 0.617\(\pm\)​0.256 0.619\(\pm\)​0.256 0.626\(\pm\)​0.272 0.619\(\pm\)​0.246 0.622\(\pm\)​0.249 0.638\(\pm\)​0.264 0.620\(\pm\)​0.251
MIMIC-III Horizon-5 0.688\(\pm\)​0.198 0.809\(\pm\)​0.260 0.729\(\pm\)​0.226 0.710\(\pm\)​0.193 0.725\(\pm\)​0.214 0.972\(\pm\)​0.373 0.694\(\pm\)​0.196
Warfarin One-step 0.036\(\pm\)​0.023 0.014\(\pm\)​0.007 0.055\(\pm\)​0.085 0.102\(\pm\)​0.109 0.270\(\pm\)​0.191 0.084\(\pm\)​0.091 0.158\(\pm\)​0.188
Warfarin Horizon-5 0.196\(\pm\)​0.143 0.152\(\pm\)​0.075 0.238\(\pm\)​0.236 0.261\(\pm\)​0.213 0.480\(\pm\)​0.315 0.749\(\pm\)​0.615 0.546\(\pm\)​0.646

3.1.0.3 Real-world factual prediction.

MIMIC-III provides a test of transfer to real-world factual ICU trajectories, where no method has access to counterfactual labels and evaluation is restricted to rolling-origin prediction under observed treatment paths. CausalLongPFN ranks first on both MIMIC-III one-step and five-step prediction. For one-step prediction, its normalized RMSE is \(0.6170\), ahead of MSM at \(0.6186\) and G-Net at \(0.6193\). For five-step prediction, it obtains \(0.6884\), ahead of GT at \(0.6938\) and G-Net at \(0.7104\). Thus, on the real clinical benchmark, the frozen synthetically pretrained model matches or exceeds domain-trained baselines without using target-domain gradient updates.

3.1.0.4 Counterfactual simulated domains.

Cancer, HIV, and warfarin provide branchable counterfactual labels, allowing direct evaluation under alternative treatment sequences. On one-step counterfactual prediction, CausalLongPFN is close to the strongest domain-trained methods: it ranks second on cancer, third on HIV, and second on warfarin. Longer-horizon performance is more mixed. CausalLongPFN remains second on warfarin and competitive on HIV, but its largest relative gap occurs on cancer five-step prediction, where RMSN, CRN, G-Net, and GT achieve lower error. This suggests that specialized recurrent or transformer models can retain an advantage when a target simulator provides enough support data for domain-specific fitting, especially over longer rollouts.

3.1.0.5 Interpretation.

Overall, the results support the main claim that broad synthetic causal pretraining can produce a useful in-context model for longitudinal treatment-response prediction. CausalLongPFN is not uniformly best, but it achieves the best domain-balanced one-step performance, the third-best domain-balanced five-step performance, and the best performance on the real MIMIC-III benchmark at both horizons. These results are notable because CausalLongPFN is evaluated as a single frozen model trained only on synthetic TSCM episodes, whereas all baselines receive domain-specific training and validation-based model selection. The pattern suggests that a sufficiently broad synthetic longitudinal causal prior can capture reusable structure across treatment-response tasks, making CausalLongPFN a strong general-purpose in-context predictor when repeated domain-specific training is expensive, rapid adaptation to a new cohort is needed, or counterfactual supervision is unavailable.

3.2 Uncertainty and Calibration↩︎

Table 3 evaluates predictive uncertainty using standard probabilistic-forecast diagnostics, including empirical coverage, NLL, CRPS, and PIT-ECE [34], [35]. Calibration varies by domain. Warfarin has the lowest RMSE, NLL, and CRPS, and its empirical coverage is slightly conservative at the \(90\%\) level. HIV also shows low point error and sharp predictive intervals, although coverage remains below nominal. MIMIC-III is the most difficult calibration setting: it has the largest NLL and CRPS and substantially wider intervals, reflecting the greater heterogeneity and noise of the real ICU benchmark. Overall, the Gaussian-mixture head provides useful distributional information without domain-specific training, but the under-coverage suggests that future work should improve uncertainty propagation, especially for multi-step prediction and real-world clinical data.

3pt

Table 3: One-step calibration of CausalLongPFN predictive distributions.Lower is better for RMSE, NLL, CRPS, and PIT-ECE. Empirical coverage shouldmatch the nominal level, while interval width should be interpreted relative tocoverage.
Domain RMSE NLL \(\downarrow\) CRPS \(\downarrow\) Pred. std. PIT-ECE \(\downarrow\) Cov. 80% Width 80% Cov. 90% Width 90%
Cancer 0.167 -0.711 0.082 0.054 0.029 0.703 0.125 0.781 0.170
HIV 0.066 -1.310 0.039 0.064 0.037 0.753 0.155 0.849 0.205
MIMIC-III 0.617 0.938 0.336 0.451 0.019 0.727 1.096 0.836 1.518
Warfarin 0.036 -1.976 0.021 0.048 0.036 0.863 0.106 0.934 0.145
Domain-balanced 0.222 -0.765 0.120 0.154 0.030 0.761 0.370 0.850 0.510

4 Conclusion↩︎

We introduced CausalLongPFN, a prior-fitted transformer for predicting history-conditional potential outcomes in longitudinal treatment-response settings. The model is pretrained only on synthetic temporal structural causal models and is then evaluated as a frozen in-context predictor on new domains. Given support trajectories, a query history, and a planned future treatment sequence, it returns a predictive distribution without target-domain gradient updates, propensity-model fitting, adversarial balancing, or simulator access at test time.

CausalLongPFN achieves the best domain-balanced one-step normalized RMSE and the third-best domain-balanced five-step normalized RMSE. It performs particularly well on factual MIMIC-III rolling-origin prediction, where it ranks first at both horizons.

These results suggest that broad synthetic causal pretraining can provide a useful in-context predictor for longitudinal treatment-response tasks, especially when retraining is costly, rapid evaluation on a new cohort is needed, or counterfactual supervision is unavailable. At the same time, the results show that domain-specific training remains valuable when sufficient target-domain data and validation signal are available.

5 Limitations and broader impact↩︎

CausalLongPFN does not remove the assumptions required for causal interpretation of longitudinal observational data. In real cohorts, counterfactual validity still depends on consistency, positivity, and sequential exchangeability given the measured history, as summarized in Appendix 6.2. Violations due to unmeasured confounding, poor treatment overlap, censoring, irregular sampling, or measurement error can bias any longitudinal counterfactual estimator, including CausalLongPFN.

The method also depends on the support of the synthetic prior. Performance may degrade when the target domain contains mechanisms, treatment policies, outcome dynamics, missingness patterns, or intervention effects that are poorly covered by the TSCM prior. The current implementation focuses on discrete treatments, fixed time grids, and deterministic mean rollout, which make the model stable, efficient, and straightforward to evaluate across heterogeneous benchmarks. These choices are not fundamental restrictions of the framework. Natural extensions include continuous or structured treatment spaces, irregular-time encoders, explicit missingness and censoring models, and stochastic rollout procedures that propagate uncertainty over future trajectories while preserving the same amortized in-context causal prediction principle.

The potential benefit of this approach is to reduce dependence on hand-built disease simulators and repeated domain-specific supervised training when studying longitudinal treatment-response prediction. A frozen in-context model could be useful for rapid benchmarking, exploratory counterfactual analysis, or settings where retraining many specialized models is impractical. The main risk is over-trust: predictions may appear precise even when the causal assumptions, data quality, treatment overlap, or prior support are inadequate. In particular, strong factual prediction on MIMIC-III should not be interpreted as validation of individual treatment effects under unobserved ICU interventions. CausalLongPFN should therefore be viewed as a research tool for causal sequence modeling and hypothesis generation, not as a standalone clinical decision system.

Code availability↩︎

Code for model training, synthetic episode generation, benchmark construction, and evaluation is available at https://github.com/Amirhossein-Zare/causal-long-pfn.

Data availability↩︎

The cancer, HIV, and warfarin benchmarks are simulated or semi-mechanistic benchmarks that can be regenerated using the released code and the simulator specifications described in the paper. MIMIC-III is a credentialed-access de-identified clinical database and is not redistributed with this paper. Reproducing MIMIC-III experiments requires obtaining access through the official data-use process and applying the preprocessing protocol described in Appendix 13.4.

Funding↩︎

No external funding was received for this work.

Competing interests↩︎

The authors declare no competing interests.

6 Causal foundations and estimand↩︎

This appendix states the longitudinal causal estimand used in the paper and the standard assumptions under which it can be interpreted causally from observational data.

6.1 Observed data and histories↩︎

For unit \(i\) at discrete time \(t\), let \[S_{i,t}\in\mathbb{R}^{d_S},\qquad A_{i,t}\in\mathcal{A},\qquad Y_{i,t}\in\mathbb{R},\qquad C_i\in\mathbb{R}^{d_C}\] denote time-varying covariates, treatment, scalar outcome, and static covariates. The model-facing longitudinal state is \[X_{i,t}=(S_{i,t},Y_{i,t})\in\mathbb{R}^d,\qquad d=d_S+1 .\] Treatment \(A_{i,t}\) is assigned after observing \(X_{i,t}\). The observed history available immediately before treatment assignment at time \(t\) is \[H_{i,t} = (C_i,X_{i,0},A_{i,0},X_{i,1},A_{i,1},\ldots,A_{i,t-1},X_{i,t}). \label{eq:history-def}\tag{4}\]

For a query unit observed through time \(t\), the model receives support trajectories \(\mathcal{C}\) from the same task or domain, the query history \(H_t\), and a planned future treatment sequence \[\bar a_{t:t+\tau-1} = (a_t,a_{t+1},\ldots,a_{t+\tau-1}).\] The target is the history-conditional potential outcome \[Y_{t+\tau}(\bar a_{t:t+\tau-1}),\] or its conditional predictive distribution given the observed query history and support trajectories: \[p\!\left( Y_{t+\tau}(\bar a_{t:t+\tau-1})\in dy \mid H_t,\mathcal{C} \right). \label{eq:appendix-target}\tag{5}\] For point prediction, we evaluate the corresponding conditional mean.

6.2 Identification assumptions↩︎

For observational data, Eq. 5 has a causal interpretation only under the usual longitudinal causal assumptions.

6.2.0.1 Consistency.

If a unit actually follows the treatment sequence \(\bar a_{t:t+\tau-1}\), then its observed outcome equals the corresponding potential outcome under that sequence.

6.2.0.2 Sequential exchangeability.

At each time point, after conditioning on the observed history \(H_t\), treatment assignment is independent of future potential outcomes. Informally, there are no unmeasured time-varying confounders after conditioning on the recorded history.

6.2.0.3 Positivity.

Every treatment sequence considered for evaluation must have positive probability, or adequate support, among units with comparable histories. Without such overlap, the corresponding counterfactual prediction requires extrapolation.

6.2.0.4 No interference and well-defined interventions.

One unit’s potential outcomes are unaffected by the treatment assignments of other units, and the treatment actions correspond to well-defined interventions.

These assumptions are standard for longitudinal causal inference and are not guaranteed by CausalLongPFN. They are required for any observational counterfactual interpretation of the predictions.

6.3 Connection to the longitudinal \(g\)-formula↩︎

Under consistency, sequential exchangeability, positivity, and no interference, the conditional mean potential outcome can be written using the longitudinal \(g\)-formula [2][4]. In words, the \(g\)-formula propagates the observed conditional transition law forward while setting future treatments to the specified intervention sequence.

Let \[K_s(dx_{s+1}\mid h_s,a_s) = \mathbb{P}(X_{s+1}\in dx_{s+1}\mid H_s=h_s,A_s=a_s)\] denote the observed one-step transition distribution. Starting from \(h_t=H_t\), define the future history recursively by appending the intervention treatment \(a_s\) and the next state \(x_{s+1}\). Then the identified conditional mean can be written schematically as \[\mathbb{E}\!\left[ Y_{t+\tau}(\bar a_{t:t+\tau-1}) \mid H_t=h_t \right] = \int y(x_{t+\tau}) \prod_{s=t}^{t+\tau-1} K_s(dx_{s+1}\mid h_s,a_s). \label{eq:appendix-gformula}\tag{6}\] This expression motivates the sequential prediction problem studied in the paper: future outcomes are predicted by repeatedly applying one-step conditional models under a specified future treatment sequence.

6.4 Role of CausalLongPFN↩︎

CausalLongPFN is an estimator for the prediction problem above. It does not introduce new identification assumptions and does not remove the need for consistency, positivity, and sequential exchangeability in observational data. Instead, it amortizes the estimation problem by pretraining on many synthetic longitudinal causal tasks and then conditioning on support trajectories from a new task at test time.

The model is trained as a one-step predictor. Multi-step predictions are obtained by deterministic plug-in rollout: the model predicts the next outcome under the planned treatment, inserts the predicted mean into the query history, and repeats this process until the desired horizon. This procedure is an approximation to full sequential predictive inference because it does not integrate over all possible intermediate outcome paths. The empirical results evaluate the resulting multi-step predictions directly in the benchmark settings.

7 Temporal structural causal prior↩︎

This appendix specifies the temporal structural causal model (TSCM) prior used to generate synthetic pretraining episodes for CausalLongPFN. Each episode draws a fresh longitudinal data-generating process from the prior and then samples support trajectories and a query trajectory from that process. The prior is intentionally heterogeneous: it varies state dimension, temporal lag structure, nonlinear mechanisms, treatment-policy confounding, latent unit heterogeneity, outcome dynamics, observation windows, and interventional rollout horizons. Its purpose is not to reproduce any single disease simulator, but to expose the model to a broad family of longitudinal treatment-response tasks with reusable causal structure.

7.1 Global ranges↩︎

Table 4: Core synthetic task ranges. The prior varies state dimension, supportsize, observation time, and prediction horizon so that a single model is trainedacross heterogeneous longitudinal causal tasks.
Quantity Value
Observed time \(t_{\mathrm{obs}}\) Uniform integer \(1\)\(60\)
Prediction horizon \(\tau\) Uniform integer \(1\)\(5\)
Maximum input length \(65\) input slots, target index up to \(65\)
State dimension \(d_S\) Uniform integer \(1\)\(10\)
Outcome dimension \(1\)
Padded input dimension \(D_{\max}\) \(11\)
Static covariate dimension \(5\); active in \(30\%\) of synthetic episodes
Treatment space \(4\) discrete treatments
Latent heterogeneity dimension \(3\)
Support size Uniform integer \(3\)\(500\)
Support anchor labels \(4\) per support trajectory
Observational query probability \(0.30\)
Support future-covariate masking probability \(0.35\)
Support target-noise augmentation \(0.15\)
Sentinel for hidden values \(-99\)

7.2 TSCM hyperparameter sampling↩︎

A synthetic TSCM instance is sampled as follows:

  1. State dimension. Sample \(d_S\sim\operatorname{Unif}\{1,\ldots,10\}\).

  2. Lag order. Sample \(K\sim\operatorname{Unif}\{1,2\}\).

  3. Instantaneous graph. Sample an instantaneous adjacency matrix \(G^{(0)}\in\{0,1\}^{d_S\times d_S}\) as a strictly lower-triangular Erdős–Rényi matrix with edge probability \[p_{\mathrm{edge}}=0.1+0.5B,\qquad B\sim\operatorname{Beta}(2,2).\] This gives an acyclic contemporaneous graph under the coordinate ordering.

  4. Lagged graph. For lag \(k\), sample a full lagged adjacency matrix \(G^{(k)}\in\{0,1\}^{d_S\times d_S}\) with edge probability \(p_{\mathrm{edge}}\gamma_{\mathrm{lag}}^k\), where \(\gamma_{\mathrm{lag}}\sim\operatorname{Unif}(0.4,0.8)\). This induces sparse temporal dependence with geometrically decaying edge probability across lags.

  5. Structural weights. Sample instantaneous weights \(W^{(0)}_{ij}\sim\mathcal{N}(0,\sigma_W^2)G^{(0)}_{ij}\) and lagged weights \(W^{(k)}_{ij}\sim\mathcal{N}(0,(0.7\sigma_W)^2)G^{(k)}_{ij}\), with \(\sigma_W\sim\operatorname{Unif}(0.3,1.0)\).

  6. Nonlinearities. Sample each generic activation independently from \[\{\mathrm{id},\tanh,\sin,\cos,|\cdot|,(\cdot)^2, \operatorname{ReLU},\operatorname{softplus}\}.\]

  7. State noise. For each state coordinate, sample a centered Gaussian, uniform, or Laplace noise family. The coordinate noise scale is zero with probability \(0.5\); otherwise it is proportional to a task-level base scale. The base scale is sampled from a low-noise range with probability \(0.6\) and from a moderate-noise range with probability \(0.4\).

  8. Autoregressive persistence. Set the coordinate-level autoregressive coefficient \(\alpha_i=0\) with probability \(0.5\); otherwise sample \(\alpha_i\sim\operatorname{Unif}(0.5,1.0)\).

  9. Treatment-policy confounding strength. Sample a policy strength multiplier that is zero with probability \(0.08\), one with probability \(0.20\), and otherwise a random integer from \(2\) to \(5\). Treatment-policy state weights are scaled by this multiplier, producing tasks with varying degrees of treatment–confounder feedback.

  10. Regime switch. With probability \(0.12\), sample a second structural mechanism with the same graph support and activate it after a sampled early-to-middle switch time. Structural weights and nonlinearities change after the switch, creating nonstationary longitudinal dynamics.

7.3 Generic structural mechanisms↩︎

For a generic non-motif state coordinate, the transition is a sparse nonlinear autoregressive update combining lagged state inputs, acyclic within-slice inputs, treatment inputs, and additive noise. Instantaneous contributions use the partially constructed next-time state \(S_{t+1,\ell}\) for \(\ell<m\), while lagged contributions use previous states from the lag buffer. Values are clipped internally to avoid numerical explosions during synthetic generation. Thus, even before adding the structured motifs below, the prior spans nonlinear autoregression, contemporaneous acyclic dependence, lagged temporal dependence, treatment effects, and heterogeneous noise.

7.4 Latent dynamical motifs↩︎

The prior optionally allocates disjoint state coordinates to five motif types. Motif coordinates are selected by a random permutation of the state dimensions, so motif identity is not tied to a fixed input channel. These motifs are included to expose the model to qualitative mechanisms common in biomedical and behavioral longitudinal data: slow accumulation, saturation, homeostatic regulation, feedback control, and proxy readout dynamics.

7.4.0.1 Action-memory channel.

With probability \(0.25\), one coordinate follows a leaky accumulation model: \[S_{t+1,m}^{\mathrm{mem}} = \delta_m S_{t,m}^{\mathrm{mem}} +w_m^\top b(A_t) +v_m^\top M_{t+1} +\varepsilon_{t+1,m},\] where \(\delta_m\sim\operatorname{Unif}(0.72,0.97)\) and \(M_{t+1}\) is a running treatment-memory vector.

7.4.0.2 Saturating channel.

With probability \(0.25\), one or two coordinates follow a nonnegative saturating update: \[S_{t+1,m}^{\mathrm{sat}} = \operatorname{clip}_{[0,6]}\!\left( S_{t,m}^{\mathrm{sat}} +r_mb_m\left(1-g_m\frac{L_t}{h_m+L_t+\epsilon}\right) -r_mS_{t,m}^{\mathrm{sat}} +\varepsilon_{t+1,m} \right),\] where \(L_t\) is a nonnegative signal derived from treatment memory and, when available, latent memory coordinates.

7.4.0.3 Homeostatic channel.

With probability \(0.25\), one coordinate reverts toward a sampled baseline: \[S_{t+1,m}^{\mathrm{hom}} = S_{t,m}^{\mathrm{hom}} +\kappa_m(\mu_m-S_{t,m}^{\mathrm{hom}}) +w_m^\top b(A_t) +\varepsilon_{t+1,m}.\]

7.4.0.4 Feedback channel.

With probability \(0.25\), one coordinate receives error-driven control from a source coordinate \(j(m)\): \[S_{t+1,m}^{\mathrm{fb}} = \rho_m S_{t,m}^{\mathrm{fb}} +\gamma_m(\eta_m-S_{t,j(m)}) +w_m^\top b(A_t) +\varepsilon_{t+1,m}.\]

7.4.0.5 Readout channel.

With probability \(0.20\), one coordinate tracks another coordinate using exponential smoothing: \[S_{t+1,m}^{\mathrm{read}} = \rho_m^{\mathrm{rd}}S_{t,m}^{\mathrm{read}} +(1-\rho_m^{\mathrm{rd}})S_{t+1,j(m)} +\varepsilon_{t+1,m}.\]

Table 5: Sampling ranges for motif-specific parameters. The motifs introduceslow accumulation, saturation, regulation, feedback, and proxy readout dynamicsinto the synthetic prior.
Motif Parameter Symbol Range
Memory decay \(\delta_m\) \([0.72,0.97]\)
Saturating baseline \(b_m\) \([0.5,1.5]\)
rate \(r_m\) \([0.02,0.15]\)
gain \(g_m\) \([0.25,0.95]\)
half-saturation \(h_m\) \([0.3,2.0]\)
Homeostatic reversion \(\kappa_m\) \([0.03,0.15]\)
baseline \(\mu_m\) \([-0.5,0.5]\)
Feedback decay \(\rho_m\) \([0.65,0.95]\)
gain \(\gamma_m\) \([0.10,0.90]\)
Readout smoothing \(\rho_m^{\mathrm{rd}}\) \([0.70,0.97]\)

7.5 Latent heterogeneity and behavior policy↩︎

Unit heterogeneity is encoded by a latent vector \(Z_i\sim\mathcal{N}(0,I_3)\) drawn once per support or query trajectory. This latent factor affects both the initial state and the treatment policy: \[\begin{align} S_{i,0} &\approx U_{S0}Z_i + \varepsilon_{i,0},\\ W_{u,i} &= W_u + U_u Z_i,\qquad u\in\{0,1\}. \end{align}\] Treatment memories used by the policy evolve as \[M_{t+1,u} = \lambda_u M_{t,u}+b_u(A_t), \qquad \lambda_u\sim\operatorname{Unif}(0.5,0.95).\] The behavior-policy logits depend on current state, recent treatment memory, static covariates when active, and latent heterogeneity. Because both baseline state and treatment assignment depend on \(Z_i\), and because treatment assignment also depends on the evolving state, support trajectories exhibit persistent unit-level heterogeneity and time-varying confounding. A small probability of near-random policy strength preserves overlap.

7.5.0.1 Synthetic treatment encoding.

In the synthetic TSCM prior, the four-valued treatment is generated through two binary policy components. Conditional on the current state, recent treatment memories, static covariates when active, and latent heterogeneity, the generator computes two logistic probabilities and samples \[A_{t,0}\sim\operatorname{Bernoulli}(p_{t,0}),\qquad A_{t,1}\sim\operatorname{Bernoulli}(p_{t,1}).\] The treatment supplied to the model is then \[A_t=A_{t,0}+2A_{t,1}\in\{0,1,2,3\}.\] This bitwise construction is used only to generate heterogeneous treatment policies. The model itself receives the resulting four-valued treatment.

Observed static covariates are included in a random subset of synthetic episodes. Specifically, the implementation activates the five-dimensional static covariate vector with probability \(0.30\) and otherwise supplies zeros. This prevents the model from assuming that static covariates are informative in every task while still exposing it to domains where baseline features are useful.

7.6 Outcome mechanism↩︎

The base readout \(R_t\) is either a selected state coordinate or an affine projection of all state variables. Coordinates belonging to dynamical motifs are sampled with elevated probability as direct outcome coordinates. The scalar outcome evolves according to an autoregressive readout with \(\rho_Y\sim\operatorname{Unif}(0.35,0.90)\), state gain in \([0.35,1.20]\), small direct and cumulative treatment effects, and weak linear trends. Outcome noise is low in most TSCMs but can be moderate in a minority of cases. Consequently, the target may depend on current state, prior outcomes, recent treatment, and accumulated treatment exposure.

7.7 Counterfactual oracle construction↩︎

For each interventional training example, the counterfactual target is constructed by structural replay:

  1. The query trajectory is simulated under the observational behavior policy from \(t=0\) to \(t_{\mathrm{obs}}\), storing the state, treatment memories, and lag buffer at \(t_{\mathrm{obs}}\).

  2. From \(t_{\mathrm{obs}}\), a second rollout is performed under the hypothetical treatment sequence \(\bar a_{t_{\mathrm{obs}}:t_\star-1}\), with future additive state noise set to its mean.

  3. The oracle outcome is continued from the observed outcome prefix \(Y_{0:t_{\mathrm{obs}}}\) using the counterfactual state path and the planned treatments, again with future outcome noise set to its mean.

This construction produces a conditional structural target given the factual query history and the planned intervention. It focuses training on the mean causal response to the supplied treatment sequence rather than on aleatoric future noise, while stochasticity remains present in support trajectories and factual query prefixes.

7.8 Support anchor time points↩︎

Each synthetic support trajectory contributes \(K_{\mathrm{sup}}=4\) labeled outcome anchors. During one-step pretraining, the first anchor is the current label time \(r+1\). The remaining anchors are the earliest post-observation label time \(t_{\mathrm{obs}}+1\), a midpoint between \(t_{\mathrm{obs}}+1\) and \(r+1\), and a random anchor sampled from \(\{t_{\mathrm{obs}}+1,\ldots,r+1\}\). This multi-anchor strategy provides in-context examples at several rollout depths from the same sampled TSCM, rather than relying on a single labeled support time point.

For external benchmark evaluation, support anchors are chosen from each support row’s available outcome times using the same four-anchor interface but a deterministic template: latest valid outcome time, midpoint, earliest valid outcome time, and one random valid anchor. Thus, the architectural interface is shared across pretraining and evaluation—multiple labeled support anchors per support trajectory—while the exact anchor-selection rule is adapted to the available benchmark rows.

7.9 Data augmentation↩︎

Three augmentations are applied during synthetic training:

  1. Observational mode with probability \(0.30\): the query target is factual under the behavior policy rather than interventional. This keeps ordinary factual prediction within the training distribution.

  2. Support target noise with probability \(0.15\): the first support target anchor may receive additive noise at scale \(0\%\), \(5\%\), or \(10\%\) of the task outcome standard deviation. This improves robustness to noisy support labels while leaving the other support anchors unchanged.

  3. Support future-covariate masking with probability \(0.35\): support state values after \(t_{\mathrm{obs}}\) and before the target horizon are replaced with the sentinel value while support outcome labels remain visible. This discourages reliance on post-intervention covariates that are unavailable for the query under hypothetical treatment sequences.

8 Training episode construction↩︎

Figure 1: Synthetic CausalLongPFN episode

8.0.0.1 Observed-prefix training and recursive evaluation.

During training, the model conditions on the query outcome history through the sampled current time \(r\) and predicts the next outcome \(Y^q_{r+1}\); all later query outcomes remain hidden. At evaluation time, only the query history through \(t_{\mathrm{obs}}\) is observed. For horizons beyond one step, the model performs plug-in sequential rollout, inserting each predicted mixture mean into the query outcome channel before predicting the next time point.

8.0.0.2 Normalization.

State normalizers use only support times \(0{:}t_{\mathrm{obs}}\): \[\mu_{S,m} = \operatorname{mean}_{j,t\le t_{\mathrm{obs}}}S_{j,t,m}, \qquad \sigma_{S,m} = \max\{\operatorname{sd}_{j,t\le t_{\mathrm{obs}}}(S_{j,t,m}),0.1\}.\] Outcome normalizers use support outcomes over \(1{:}t_\star\): \[\mu_Y = \operatorname{mean}_{j,1\le t\le t_\star}Y_{j,t}, \qquad \sigma_Y = \max\{\operatorname{sd}_{j,1\le t\le t_\star}(Y_{j,t}),0.1\}.\] Episodes with near-constant outcome scale are rejected. Normalized state values are clipped to \([-3,3]\), normalized outcomes to \([-10,10]\), and unavailable values are marked with the sentinel \(-99\).

9 Model architecture details↩︎

9.1 State encoder↩︎

Let \(V_t\in\mathbb{R}^{D_{\max}}\) denote the padded model input at time \(t\), containing the time-varying covariates and scalar outcome, and let \(m_t=\mathbb{I}\{V_t<-90\}\) denote the hidden-value mask induced by the sentinel. Hidden entries are set to zero before projection, while the mask itself is retained as an input feature. First differences are scaled by \(0.5\) and set to zero whenever either adjacent value is hidden.

Let \(y\) denote the active outcome coordinate. The encoder separates covariate and outcome channels: \[\begin{align} e_t^{S} &= W_S[V_t^{(-y)},0.5\Delta V_t^{(-y)},-2m_t^{(-y)}],\\ e_t^{Y} &= W_Y[V_t^{(y)},0.5\Delta V_t^{(y)},-2m_t^{(y)}],\\ e_t^{A} &= W_A\operatorname{onehot}(A_t). \end{align}\] The encoded timestep representation is \[e_t=\operatorname{LN}(e_t^S+e_t^Y+e_t^A).\] Separating the outcome channel from the covariate channels helps preserve the distinction between observed predictors and the target process. Padded inactive dimensions remain zero after cleaning, and the input scale factor \(\sqrt{D_{\max}/d}\) helps keep signal magnitudes comparable across tasks with different active state dimensions.

9.2 History encoder↩︎

The history encoder is a causal transformer that maps each longitudinal trajectory to time-indexed history representations. It uses:

  • \(4\) layers of pre-norm self-attention;

  • model dimension \(256\), \(8\) attention heads, and feedforward dimension \(1024\);

  • sinusoidal temporal positional encodings up to the maximum sequence length;

  • a causal attention mask preventing each time point from attending to future positions;

  • zero initialization of selected residual output projections for stable training from scratch.

For a query current time \(r\), the representation \(h_r\) is extracted at position \(r\). For a support anchor label \(Y_s\), the corresponding history representation is extracted at \(s-1\), so the support token represents the predictive mapping from history and treatment through time \(s-1\) to the outcome label at time \(s\).

9.3 PFN context encoder↩︎

The PFN context encoder performs in-context adaptation over support trajectories. All \(n_{\mathrm{ctx}}K_{\mathrm{sup}}\) support-anchor tokens and the single query token attend bidirectionally to each other using full self-attention, with a key-padding mask for padded support slots. No positional encoding is added for the arbitrary order of support trajectories. Each PFN layer uses multi-head attention, GELU feedforward blocks, residual connections, layer normalization, and zero-initialized residual output projections.

For support trajectory \(j\) and anchor time \(s_{jk}\), the support token is \[z^{\mathrm{ctx}}_{jk} = W_{\mathrm{tok}}\left[ h_{j,s_{jk}-1} +W_xV_{j,s_{jk}-1} +W_CC_j +W_Gg(\mathcal{C}),\; W_y(y_{j,s_{jk}},0) \right], \label{eq:support95token}\tag{7}\] where \(g(\mathcal{C})\) contains symmetric support-level outcome statistics, including the mean and standard deviation of support anchor outcomes. The query token at current time \(r\) is \[z^{\mathrm{qry}}_{r} = W_{\mathrm{tok}}\left[ h^q_r +W_xV^q_r +W_CC_q +W_Gg(\mathcal{C}),\; e_{\mathrm{qry}} \right], \label{eq:query95token}\tag{8}\] where \(e_{\mathrm{qry}}\) is a learned query-label embedding. Thus, support tokens contain observed anchor labels, whereas the query token marks the unknown target to be predicted.

9.4 Gaussian mixture head↩︎

The final query representation \(u_r\) is mapped to the parameters of a five-component Gaussian mixture. The mixture weights, residual means, and scales are computed as \[\begin{align} \log\pi_r &= \log\operatorname{softmax}(W_\pi u_r/T_\pi), \qquad T_\pi=1.0,\\ \Delta\mu_r &= 7\tanh(W_\mu u_r/7),\\ \sigma_r &= \operatorname{clip}_{[0.02,2.0]} \{\operatorname{softplus}(W_\sigma u_r)+0.02\}. \end{align}\] The component means are residualized around the most recent visible or self-predicted query outcome: \[\mu_{r,k} = \operatorname{clip}_{[-12,12]}(y^q_r+\Delta\mu_{r,k}).\] The mean projection is initialized at zero, producing an initial persistence predictor. This initialization stabilizes early training because \(\mu_{r,k}\approx y^q_r\) before the model has learned task-specific dynamics.

9.5 Architecture and prior hyperparameters↩︎

Table 6: Architecture and synthetic-prior hyperparameters. Training andoptimization settings are reported separately inTable 7.
Hyperparameter Symbol Value
Architecture
Model dimension \(d_{\mathrm{model}}\) \(256\)
Attention heads \(h\) \(8\)
History encoder layers \(N_{\mathrm{enc}}\) \(4\)
PFN layers \(N_{\mathrm{pfn}}\) \(6\)
Feedforward dimension \(d_{\mathrm{ff}}\) \(1024\)
Dropout \(0.10\)
GMM components \(K_{\mathrm{gmm}}\) \(5\)
Mixture temperature \(T_\pi\) \(1.0\)
Minimum / maximum GMM std. dev. \(0.02\) / \(2.0\)
Residual mean bound before clipping \([-7,7]\)
Final mean clipping \([-12,12]\)
Synthetic prior
Maximum state dimension \(d_{S,\max}\) \(10\)
Observation window \(t_{\mathrm{obs}}\) \(1\)\(60\)
Prediction horizon \(\tau\) \(1\)\(5\)
Support size \(n_{\mathrm{ctx}}\) \(3\)\(500\)
Number of treatments \(|\mathcal{A}|\) \(4\)
Latent unit dimension \(3\)
Support anchors \(K_{\mathrm{sup}}\) \(4\)

9.5.0.1 Initialization.

The self-attention output projections and final feed-forward projections in the history and PFN transformer blocks are initialized at zero. The final GMM mean projection is also initialized at zero, producing an initial persistence forecast \(\mu_{r,k}\approx y_r\).

10 Loss function details↩︎

The model is trained with a Gaussian-mixture one-step predictive loss. The total loss is \[\mathcal{L} = \mathcal{L}_{\mathrm{NLL}} +\lambda_m\mathcal{L}_{\mathrm{mean}} +\lambda_c\mathcal{L}_{\mathrm{conc}}, \qquad \lambda_m=0.25,\quad \lambda_c=0.03.\]

10.0.0.1 Robust NLL.

For normalized target \(z\), the Gaussian-mixture negative log-likelihood is \[\ell_{\mathrm{NLL}} = -\log \sum_{k=1}^{K_{\mathrm{gmm}}} \pi_k \frac{1}{\sigma_k\sqrt{2\pi}} \exp\left\{ -\frac{1}{2} \left(\frac{z-\mu_k}{\sigma_k}\right)^2 \right\}.\] The implementation uses a linear tail for very large NLL values: \[\tilde{\ell}_{\mathrm{NLL}} = \begin{cases} \ell_{\mathrm{NLL}}, & \ell_{\mathrm{NLL}}\le 15,\\ 15+0.01(\ell_{\mathrm{NLL}}-15), & \ell_{\mathrm{NLL}}>15. \end{cases}\] This robustification prevents rare unstable synthetic examples from dominating gradients while retaining a loss signal.

10.0.0.2 Mean loss.

The predictive mean is \[\hat{z}=\sum_k\pi_k\mu_k.\] The auxiliary mean loss is \[\mathcal{L}_{\mathrm{mean}} = \operatorname{Huber}_{\delta=3}(\hat{z}-z).\] This term provides a direct gradient signal for point prediction and stabilizes early optimization.

10.0.0.3 Concentration penalty.

The mixture-concentration penalty is \[\mathcal{L}_{\mathrm{conc}} = \left[\max_k\pi_k-0.90\right]_+.\] It discourages premature collapse of the mixture distribution onto a single component.

11 Training procedure and stability↩︎

11.0.0.1 Optimizer and schedule.

The model is optimized with AdamW [36]. Weight decay is applied to ordinary dense weight matrices, while biases, layer-normalization parameters, query embeddings, context-statistic encoders, and static-feature encoders are excluded from weight decay. The learning rate is warmed up linearly for \(400\) optimizer steps and then cosine-decayed to \(2\%\) of the peak over \(10000\) steps.

11.0.0.2 Stochastic PFN depth.

At each training step, the number of active PFN layers is sampled uniformly from \(\{3,\ldots,6\}\). This stochastic-depth-like regularization encourages useful intermediate-depth representations and reduces dependence on the deepest PFN stack.

11.0.0.3 Gradient accumulation and clipping.

Gradients are accumulated over \(16\) micro-batches, giving an effective batch size of \(256\) synthetic episodes. Before each optimizer step, gradients are unscaled under automatic mixed precision and clipped using a threshold that increases linearly from \(0.5\) to \(1.5\) over the first \(4000\) optimizer steps. This applies tighter clipping during early training and looser clipping after the model stabilizes.

11.0.0.4 Numerical stability safeguards.

Training includes several numerical safeguards:

  1. If the loss is non-finite, the batch is skipped and the AMP loss scale is reduced.

  2. If the global gradient norm is non-finite, the optimizer step is skipped.

  3. The GMM loss computation is upcast to FP32 before log-sum-exp operations.

  4. Normalized inputs and targets are clipped, and GMM standard deviations are bounded in \([0.02,2.0]\).

Table 7: Training and optimization hyperparameters. Gradient accumulation givesan effective batch size of \(256\) synthetic episodes.
Quantity Value
Batch size \(16\)
Gradient accumulation \(16\) steps
Effective batch size \(256\)
Optimizer AdamW
Learning rate \(3\times10^{-4}\)
Weight decay \(10^{-5}\), excluding bias/norm/static/query parameters
Warmup \(400\) optimizer steps
Schedule Cosine decay to \(0.02\) of base LR over \(10000\) steps
Maximum optimizer steps \(10000\)
Random PFN depth during training Uniformly \(3\)\(6\) layers
Gradient clipping Linear ramp \(0.5\) to \(1.5\) over \(4000\) steps
Checkpoint interval \(500\) optimizer steps
Mixed precision Enabled on CUDA
Seed \(42\)

12 Autoregressive rollout↩︎

Figure 2: Plug-in autoregressive counterfactual rollout

The final mixture is conditional on the self-fed mean trajectory. Thus the reported distribution does not integrate over all possible intermediate outcome paths. A stochastic ancestral variant could sample from the mixture at each intermediate step and repeat the rollout to approximate full path-level uncertainty.

13 Evaluation datasets↩︎

We evaluate on four longitudinal treatment-response benchmarks: a cancer tumor growth simulator, a semi-mechanistic warfarin pharmacokinetic/pharmacodynamic (PK/PD) simulator, an HIV ODE simulator based on Adams/WhyNot dynamics, and a factual MIMIC-III ICU rolling-origin benchmark. Cancer, warfarin, and HIV are branchable simulated or semi-mechanistic systems: for these domains, counterfactual outcomes under alternative future treatment sequences are available by replaying the same patient-specific dynamics under intervened treatments. MIMIC-III is real observational ICU data and does not reveal outcomes under unobserved interventions; it is therefore used only for factual rolling-origin prediction under observed future treatments. This distinction follows prior longitudinal counterfactual evaluations, where simulated systems provide ground-truth counterfactual labels and real ICU cohorts provide factual temporal prediction benchmarks [5][8].

13.0.0.1 Common task construction.

Let \(i\) index patients and let \(t\) index discrete time. We write \(S_{i,t}\) for time-varying covariates or simulator state variables, \(A_{i,t}\) for treatment, \(Y_{i,t}\) for the scalar target outcome, and \(C_i\) for time-invariant patient features. The model-facing state is \(X_{i,t}=(S_{i,t},Y_{i,t})\) when the outcome channel is included. Each rolling-origin query observes a patient history up to an origin time and asks for either the next outcome or the outcome after a supplied future treatment sequence.

Across domains, raw trajectories have length \(T=60\), the projection horizon is \(H=5\), and the configured minimum observed history length is \(t_{\min}=10\). For each domain, tasks vary the confounding level \(\gamma\), support size \(n_{\mathrm{sup}}\in\{40,80,160,320,500\}\), and random repetition. In the simulated and semi-mechanistic domains, \(\gamma\) controls how strongly the behavior policy depends on current patient state and therefore controls the strength of time-varying confounding. In MIMIC-III, \(\gamma\) is retained only as a stratification variable for consistent task organization and does not modify the observed data.

All reported metrics use support-only normalization. The rolling-origin filters, indexing conventions, and clipping rules used for scoring are given in Appendix 15. The domain-level summary is shown in Table 8.

13.0.0.2 Counterfactual and factual test queries.

For cancer, warfarin, and HIV, one-step query rows branch from a factual patient state and evaluate alternative treatments. Horizon-5 query rows branch from a factual origin and replay the same patient-specific dynamics forward under randomly sampled treatment sequences of length \(H=5\). For MIMIC-III, both one-step and horizon-5 rows are factual rolling-origin predictions: future treatments are the observed ICU treatments, not interventions. In all domains, scored rows follow the common evaluation protocol in Appendix 15.

Table 8: Evaluation datasets. Simulated and semi-mechanistic domains providebranchable counterfactual labels; MIMIC-III is factual-only and evaluatestemporal prediction under observed clinical practice.
Domain Data type Treatments Target Query construction
Cancer tumor growth Fully simulated PK/PD tumor dynamics Four discrete treatment actions induced by chemotherapy/radiotherapy combinations \(\log(1+\text{clipped tumor volume})\) One-step: all four joint actions. Multi-step: random treatment sequences of length \(H\).
Warfarin Semi-mechanistic PK/PD simulator Four dose classes corresponding to \(0,2,5,10\) mg/day, delivered every 4-hour bin INR One-step: all four dose actions. Multi-step: random dose sequences of length \(H\).
HIV Adams/WhyNot-style six-compartment ODE simulator Four antiretroviral regimens: none, PI only, RTI only, RTI+PI \(\log_{10}(1+\text{free virus})\) One-step: all four regimens. Multi-step: random regimen sequences of length \(H\).
MIMIC-III Real factual ICU time series from MIMIC-Extract Four observed treatment classes from vasopressors and ventilation: none, vaso, vent, vaso+vent Diastolic blood pressure Factual-only rolling-origin rows; future treatments are observed ICU treatments, not interventions.

All targets are normalized only at scoring time using support-set statistics, as described in Appendix 15.

13.1 Cancer tumor growth simulator↩︎

13.1.0.1 Background.

The cancer benchmark follows the tumor-growth simulator used in RMSN, CRN, Causal Transformer, G-Net, and related longitudinal counterfactual evaluations [5][8], [23]. The simulator represents non-small-cell lung cancer tumor volume evolving under chemotherapy and radiotherapy. It combines Gompertz-style tumor growth with linear-quadratic radiotherapy effects and log-cell-kill chemotherapy effects. Treatment assignment depends on recent tumor history, producing time-varying confounding.

13.1.0.2 Dynamics.

Let \(Y_{i,t}^{\mathrm{raw}}\) denote raw tumor volume. Diameter and volume are converted by \[\operatorname{Vol}(d) = \frac{4}{3}\pi\left(\frac{d}{2}\right)^3, \qquad \operatorname{Diam}(y) = 2\left(\frac{y}{4\pi/3}\right)^{1/3}.\] The carrying capacity is \(K=\operatorname{Vol}(30)\), and the death threshold is \(Y_{\max}=\operatorname{Vol}(13)\). Patient-specific parameters include tumor growth rate \(\rho_i\), radiosensitivity coefficients \(\alpha_i\) and \(\beta_i=\alpha_i/10\), and chemotherapy kill coefficient \(\beta_i^c\).

At each time \(t\), chemotherapy and radiotherapy are assigned by Bernoulli policies depending on recent tumor diameter: \[\bar D_{i,t}^{(15)} = \frac{1}{|\mathcal{W}_t|} \sum_{s\in\mathcal{W}_t} \operatorname{Diam}\!\left(Y_{i,s}^{\mathrm{raw}}\right), \qquad \mathcal{W}_t=\{\max(0,t-15),\ldots,t\}.\] The behavior policy is \[\Pr(A^c_{i,t}=1\mid \bar H_{i,t}) = \Pr(A^r_{i,t}=1\mid \bar H_{i,t}) = \sigma\!\left[ \frac{\gamma}{D_{\max}} \left( \bar D_{i,t}^{(15)}-\frac{D_{\max}}{2} \right) \right],\] where \(A^c_{i,t}\) and \(A^r_{i,t}\) denote chemotherapy and radiotherapy indicators and \(D_{\max}=13\). Larger \(\gamma\) strengthens the dependence of treatment assignment on tumor history.

Chemotherapy is administered as a dose of 5 units when \(A^c_{i,t}=1\), with half-life one time step: \[C_{i,t}=2^{-1}C_{i,t-1}+5A^c_{i,t}.\] Radiotherapy is an immediate dose \(R_{i,t}=2A^r_{i,t}\). Raw tumor volume evolves as \[\begin{align} Y_{i,t+1}^{\mathrm{raw}} = Y_{i,t}^{\mathrm{raw}} \bigg[ 1 +\rho_i\log\!\left(\frac{K}{Y_{i,t}^{\mathrm{raw}}}\right) -\beta^c_i C_{i,t} -\left(\alpha_i R_{i,t}+\beta_i R_{i,t}^2\right) +\epsilon_{i,t} \bigg], \qquad \epsilon_{i,t}\sim\mathcal{N}(0,0.01^2). \end{align} \label{eq:cancer95dyn}\tag{9}\]

13.1.0.3 Outcome representation.

The model-facing cancer outcome is clipped log-volume, \[Y_{i,t} = \log\!\left(1+\min\{Y_{i,t}^{\mathrm{raw}},Y_{\max}\}\right).\] Support and query outcomes are normalized from this transformed scale.

13.2 Warfarin semi-mechanistic PK/PD simulator↩︎

13.2.0.1 Background.

The warfarin benchmark is a semi-mechanistic PK/PD simulator motivated by standard warfarin dose–response models: oral absorption and elimination, delayed anticoagulant response through inhibition of vitamin-K-dependent coagulation-factor synthesis, INR readout from clotting-factor activity, and patient heterogeneity driven by CYP2C9 metabolism, VKORC1 sensitivity, age, dietary vitamin K, and adherence [24][26].

13.2.0.2 PK model.

Each time step is a 4-hour bin. The treatment space is \(\mathcal{A}=\{0,1,2,3\}\), corresponding to daily dose levels \[(0,2,5,10)\;\mathrm{mg/day}.\] The PK state follows a gut \(\rightarrow\) plasma \(\rightarrow\) effect-site model: \[\begin{align} \dot{A}_g &= -k_a A_g,\\ \dot{C}_p &= k_a A_g/V_d - k_e C_p,\\ \dot{C}_e &= k_{e0}(C_p-C_e), \end{align}\] where \(A_g\) is gut depot amount, \(C_p\) is plasma concentration, \(C_e\) is effect-site concentration, \(k_a\) is absorption, \(k_e=\mathrm{CL}/V_d\) is elimination, and \(k_{e0}\) is effect-site equilibration.

13.2.0.3 PD model.

Effect-site concentration inhibits vitamin-K-dependent synthesis through an \(E_{\max}\) model: \[I(t) = E_{\max}\frac{C_e(t)^h}{C_e(t)^h+\mathrm{EC}_{50}^h}.\] Each coagulation factor \(f\in\{\mathrm{II},\mathrm{VII},\mathrm{X},\mathrm{PC}\}\) follows delayed turnover: \[\dot{f}(t) = k_{\mathrm{out},f} \left[ \mathrm{VK}(t)\{1-s_f I(t)\}-f(t) \right].\] INR is computed from factor deficits: \[\begin{align} \Delta_f(t) &= 1.30[1-f_{\mathrm{VII}}(t)]_+ +0.95[1-f_{\mathrm{X}}(t)]_+ +0.80[1-f_{\mathrm{II}}(t)]_+,\\ \mathrm{PT}(t) &= 1+1.6\Delta_f(t)+0.70\Delta_f(t)^2,\\ \mathrm{INR}(t) &= \mathrm{INR}_0\,\mathrm{PT}(t)^{\mathrm{ISI}}. \end{align}\]

13.2.0.4 Patient heterogeneity and confounding.

Patient heterogeneity includes CYP metabolism class, VKORC1 sensitivity, age, clearance, absorption, effect-site kinetics, pharmacodynamic sensitivity, vitamin-K baseline, clinic bias, and maintenance-dose requirement. The behavior policy is a softmax over dose classes whose logits depend on current INR, distance from the therapeutic range, INR trend, effect-site concentration, recent dose load, adherence, age, maintenance-dose class, and clinic bias. The confounding parameter \(\gamma\) scales the INR-dependent policy terms, increasing the dependence of dosing on current patient state.

13.2.0.5 State and outcome.

The visible model-facing state is 10-dimensional: \[\begin{align} X_t = \big( &C^{\mathrm{plasma}}_t,\, C^{\mathrm{effect}}_t,\, F^{\mathrm{II}}_t,\, F^{\mathrm{VII}}_t,\, K^{\mathrm{vit}}_t,\, \mathrm{INR}_t,\\ &\mathrm{doseLoad}_{7d,t},\, \mathrm{CYP}_i,\, \mathrm{VKORC1}_i,\, \mathrm{ageNorm}_i \big). \end{align}\] The scalar target is \[Y_t=\mathrm{INR}_t.\]

13.3 HIV Adams/WhyNot ODE simulator↩︎

13.3.0.1 Background.

The HIV benchmark is based on the six-compartment Adams HIV treatment ODE used in the WhyNot simulator suite [27], [28]. It models immunological and virological dynamics under antiretroviral therapy and allows patient-specific counterfactual evaluation by intervening on future treatment regimens.

13.3.0.2 ODE dynamics.

The raw state is \[S_t=(T_{1,t},T^*_{1,t},T_{2,t},T^*_{2,t},V_t,E_t),\] where \(T_1,T_2\) are uninfected target-cell populations, \(T^*_1,T^*_2\) are infected cell populations, \(V\) is free virus, and \(E\) is immune response. The dynamics follow \[\begin{align} \dot{T}_1 &= \lambda_1-d_1T_1-(1-\epsilon_1)k_1VT_1,\\ \dot{T}_1^* &= (1-\epsilon_1)k_1VT_1-\delta T_1^*-m_1ET_1^*,\\ \dot{T}_2 &= \lambda_2-d_2T_2-(1-f\epsilon_1)k_2VT_2,\\ \dot{T}_2^* &= (1-f\epsilon_1)k_2VT_2-\delta T_2^*-m_2ET_2^*,\\ \dot{V} &= (1-\epsilon_2)N_T\delta(T_1^*+T_2^*)-cV -\left[ (1-\epsilon_1)\rho_1k_1T_1 +(1-f\epsilon_1)\rho_2k_2T_2 \right]V,\\ \dot{E} &= \lambda_E +\frac{b_E(T_1^*+T_2^*)}{T_1^*+T_2^*+K_B}E -\frac{d_E(T_1^*+T_2^*)}{T_1^*+T_2^*+K_D}E -\delta_EE. \end{align}\]

13.3.0.3 Treatment space.

There are four treatment regimens: \[\begin{array}{ccl} 0 &:& \text{no therapy},\quad (\epsilon_1,\epsilon_2)=(0,0),\\ 1 &:& \text{PI only},\quad (\epsilon_1,\epsilon_2)=(0,0.3),\\ 2 &:& \text{RTI only},\quad (\epsilon_1,\epsilon_2)=(0.7,0),\\ 3 &:& \text{RTI+PI},\quad (\epsilon_1,\epsilon_2)=(0.7,0.3). \end{array}\]

13.3.0.4 Patient heterogeneity and confounding.

Patient heterogeneity is introduced through perturbations of the ODE parameters, individual RTI/PI efficacy scales, viral and immune thresholds, and policy aggressiveness. The behavior policy computes a severity score from current \(\log_{10}(V+1)\) and \(\log_{10}(E+1)\), scales this score by \(\gamma\), adds inertia for the previous treatment, and samples one of the four regimens from a softmax. Larger \(\gamma\) increases dependence of treatment choice on biological state.

13.3.0.5 State and outcome.

The model-facing state is the log-transformed compartment vector \[X_{t,k}=\log_{10}(1+S_{t,k}),\qquad k=1,\ldots,6.\] The scalar target is transformed free virus, \[Y_t=\log_{10}(1+V_t).\]

13.4 MIMIC-III factual ICU rolling-origin benchmark↩︎

13.4.0.1 Background.

The MIMIC benchmark is constructed from MIMIC-III ICU stays using a MIMIC-Extract-style hourly representation [29][33]. Because MIMIC-III does not provide ground-truth counterfactual outcomes, we use it only for factual rolling-origin prediction under observed future treatments.

13.4.0.2 State and preprocessing.

Each ICU stay is treated as a single hourly sequence. The model-facing state is 10-dimensional: \[\begin{align} X_t = (& \text{diastolic blood pressure}, \text{mean blood pressure}, \text{oxygen saturation}, \text{heart rate}, \text{respiratory rate},\\ & \text{Glasgow Coma Scale total}, \text{glucose}, \text{creatinine}, \text{bicarbonate}, \text{sodium} ). \end{align}\] Static features are derived from demographic variables and represented by a fixed-dimensional vector.

13.4.0.3 Treatment space and outcome.

The two binary treatment indicators are vasopressor administration \(\mathrm{vaso}_t\) and mechanical ventilation \(\mathrm{vent}_t\). They are combined into the four-valued treatment used by the shared model interface, \[A_t = \mathbb{I}\{\mathrm{vaso}_t\} +2\mathbb{I}\{\mathrm{vent}_t\},\] with mapping \[0:\text{none},\qquad 1:\text{vaso},\qquad 2:\text{vent},\qquad 3:\text{vaso+vent}.\] The scalar target is \[Y_t=\text{diastolic blood pressure}_t.\] MIMIC-III results should be interpreted as factual temporal prediction results, not as validation of counterfactual treatment effects.

14 Baseline Models↩︎

We compare against six longitudinal baselines: a classical Marginal Structural Model (MSM), Recurrent Marginal Structural Networks (RMSN), G-Net, Counterfactual Recurrent Networks (CRN), Causal Transformer (CT), and G-Transformer (GT). Together, these baselines cover the main adjustment strategies used in longitudinal treatment-response prediction: inverse-probability weighting [3], recurrent marginal structural modeling [5], neural \(g\)-computation [7], adversarial representation balancing [6], and transformer-based counterfactual sequence modeling [8], [9].

14.0.0.1 Common notation.

For unit \(i\) and time \(t\), let \(x_{i,t}\in\mathbb{R}^{d_x}\) denote the baseline covariate input corresponding to the time-varying covariates \(S_{i,t}\), let \(y_{i,t}\in\mathbb{R}\) denote the scalar outcome \(Y_{i,t}\), let \(a_{i,t}\in\{0,1,2,3\}\) denote the treatment \(A_{i,t}\), and let \(c_i\in\mathbb{R}^{5}\) denote static covariates corresponding to \(C_i\). We write \(\tilde{y}_{i,t}\) for the normalized outcome used by the baseline training code. All baselines follow the same inclusive observation-time convention as Section 2.1: the model observes the history through \(t_{\mathrm{obs}}\), receives the first planned treatment \(a_{i,t_{\mathrm{obs}}}\), and predicts future outcomes under the supplied treatment sequence.

14.0.0.2 Normalization and metrics.

Continuous inputs are normalized using statistics computed from the support trajectories of the corresponding benchmark file. The primary reported metric is normalized RMSE under the shared support-only evaluation normalization. For CausalLongPFN, predictions are first converted from the model’s internal PFN-context normalization to raw outcome units and then to the shared evaluation normalization. Full scoring details, including clipping rules, are given in Appendix 15.

14.0.0.3 Treatment encodings.

All methods use the same four-valued treatment space \(a_t\in\{0,1,2,3\}\). For models that require vector-valued treatment inputs, we use either the one-hot encoding \[\phi_4(a_t)=e_{a_t}\in\{0,1\}^{4},\] or the equivalent two-bit decomposition \[\phi_2(a_t) = \left(a_t \bmod 2,\; \left\lfloor a_t/2\right\rfloor\right) \in\{0,1\}^{2}.\] CausalLongPFN and the neural baselines use the four-valued treatment input, whereas MSM and RMSN use \(\phi_2\) in their propensity-weighting components.

14.0.0.4 Hyperparameter tuning protocol.

Baseline hyperparameters are selected using only support trajectories from the target domain. For each baseline, the evaluation runner performs an initial random search over the method-specific search space in Table 9 for the first dataset in each \((\text{domain}, n_{\mathrm{sup}})\) tuning group. Candidate configurations are trained on a support-training split and ranked using normalized RMSE on a held-out support-validation split. The top cached candidate is then reused and re-evaluated on subsequent support-validation splits within the same tuning group. The selected configuration is finally refit on the full support set before query evaluation. Query outcomes are never used for hyperparameter selection.

This protocol gives all baselines domain-specific supervision and validation-based model selection. In contrast, CausalLongPFN is evaluated as a frozen pretrained model: no target-domain gradients are taken, no validation set is used for model selection, and no target-domain hyperparameters are tuned.

3pt

Table 9: Baseline hyperparameter search spaces used by the evaluation runner.Each baseline is tuned using support-set validation only, with random searchover the listed discrete candidates. CausalLongPFN is not included because itis evaluated frozen without target-domain tuning.
Method Architecture / model search space Optimization / regularization search space
MSM Regressor \(\in\{\mathrm{linear},\mathrm{ridge}\}\); lag features \(=2\); ridge penalty \(\alpha\in\{0.1,1.0,10.0\}\) Stabilized treatment weights are clipped at support-set quantiles \((0.01,0.99)\). Logistic propensity models use maximum iteration count \(500\) and \(C=10^6\).
RMSN Number of recurrent layers \(\in\{1,2\}\). Encoder and decoder hidden widths are selected from a data-dimensionality-aware grid. Let \(d_x\) be the number of time-varying covariates, \(C_{\mathrm{hist}}=d_x+1+2+5\), and \(C_{\mathrm{dec}}=1+2+5\). Candidate widths are obtained by multiplying \(C_{\mathrm{hist}}\) and \(C_{\mathrm{dec}}\) by \(\{0.5,1,2,4\}\), rounding to a multiple of \(16\), clipping to \([32,160]\), and unioning with \(\{32,48,64,96,128\}\). Dropout \(\in\{0.1,0.2,0.3,0.4,0.5\}\); propensity learning rate \(\in\{10^{-2},10^{-3},10^{-4}\}\); encoder learning rate \(\in\{10^{-2},10^{-3},10^{-4}\}\); decoder learning rate \(\in\{10^{-2},10^{-3},10^{-4}\}\); encoder batch size \(\in\{64,128,256\}\); decoder batch size \(\in\{256,512,1024\}\); gradient clipping \(\in\{0.5,1.0,2.0,4.0\}\).
G-Net Hidden size \(\in\{48,64,96,128\}\); representation size \(\in\{48,64,96\}\); number of recurrent layers \(\in\{1,2\}\). Dropout \(\in\{0.05,0.10,0.20\}\); learning rate \(\in\{10^{-3},3{\times}10^{-4}\}\); batch size \(\in\{32,64\}\); epochs \(\in\{80,120\}\); covariate/vitals loss weight \(\in\{0.15,0.25,0.35\}\).
CRN Number of recurrent layers \(\in\{1,2\}\). Hidden, balanced-representation, and fully connected widths are selected from a data-dimensionality-aware grid. Let \(C_{\mathrm{hist}}=4+d_x+1+5\) and \(C_{\mathrm{dec}}=4+1+5\). Width candidates are obtained by multiplying these quantities by \(\{0.5,1,2,4\}\), rounding to a multiple of \(16\), clipping to \([32,160]\), and unioning hidden and balanced widths with \(\{32,48,64,96,128\}\) and fully connected widths with \(\{32,64,96,128,192\}\). Dropout \(\in\{0.1,0.2,0.3,0.4,0.5\}\); encoder learning rate \(\in\{10^{-2},10^{-3},10^{-4}\}\); decoder learning rate \(\in\{10^{-2},10^{-3},10^{-4}\}\); encoder batch size \(\in\{64,128,256\}\); decoder batch size \(\in\{256,512,1024\}\); gradient clipping \(\in\{0.5,1.0,2.0\}\).
Causal Transformer Transformer layers \(\in\{2,3,4\}\); attention heads \(\in\{2,4\}\); sequence hidden size \(\in\{64,96,128\}\); balanced-representation size \(\in\{64,96,128\}\); fully connected hidden size \(\in\{64,96,128\}\). Dropout \(\in\{0.05,0.10,0.20\}\); learning rate \(\in\{10^{-3},3{\times}10^{-4},10^{-4}\}\); weight decay \(\in\{10^{-5},10^{-4},10^{-3}\}\); batch size \(\in\{32,64\}\); gradient clipping \(\in\{0.5,1.0\}\); treatment loss weight \(\in\{0.05,0.10,0.20\}\).
G-Transformer Transformer layers \(\in\{2,3,4\}\); attention heads \(\in\{2,4\}\); model dimension \(\in\{32,48,64,96,128\}\); balanced-representation size \(\in\{32,48,64,96,128\}\); fully connected hidden size \(\in\{32,64,96,128,192\}\). Dropout \(\in\{0.1,0.2,0.3\}\); learning rate \(\in\{10^{-3},10^{-4},10^{-5}\}\); weight decay \(\in\{10^{-5},10^{-4},10^{-3}\}\); batch size \(\in\{16,32,64\}\).

14.1 Marginal Structural Model↩︎

The MSM baseline adjusts for time-varying confounding through stabilized inverse probability of treatment weights [3]. The numerator propensity model uses prior treatment history, \[z^{\mathrm{num}}_{i,t} = \sum_{u=0}^{t-1}\phi_2(a_{i,u}),\] whereas the denominator propensity model conditions on treatment, covariate, outcome, and static history, \[z^{\mathrm{den}}_{i,t} = \left[ \sum_{u=0}^{t-1}\phi_2(a_{i,u}),\; x_{i,t-L_{\mathrm{lag}}:t},\; \tilde{y}_{i,t-L_{\mathrm{lag}}:t},\; c_i \right].\] The stabilized treatment ratio at time \(t\) is \[w_{i,t} = \frac{ \prod_{b=1}^{2} p_{\mathrm{num},b} \left(a_{i,t,b}\mid z^{\mathrm{num}}_{i,t}\right) }{ \prod_{b=1}^{2} p_{\mathrm{den},b} \left(a_{i,t,b}\mid z^{\mathrm{den}}_{i,t}\right) },\] where \(a_{i,t,b}\) is the \(b\)th treatment bit. The outcome model is direct in horizon: for horizon \(\tau\), it predicts \(\tilde{y}_{i,t+\tau}\) from \[z^{\mathrm{MSM}}_{i,t,\tau} = \left[ z^{\mathrm{den}}_{i,t},\; \sum_{u=t}^{t+\tau-1}\phi_2(a_{i,u}) \right].\]

14.2 Recurrent Marginal Structural Networks↩︎

RMSN replaces the propensity and outcome regressions of MSM with recurrent neural networks [5]. A treatment-only propensity network predicts the current treatment from previous treatments, \[\hat{p}^{\mathrm{num}}_{i,t} = \sigma\!\left(f_{\mathrm{prop},T}(\phi_2(a_{i,0:t-1}))\right),\] while a history-dependent propensity network predicts treatment from previous treatments, covariates, outcomes, and static features, \[\hat{p}^{\mathrm{den}}_{i,t} = \sigma\!\left( f_{\mathrm{prop},H}(\phi_2(a_{i,0:t-1}),x_{i,0:t}, \tilde{y}_{i,0:t},c_i) \right).\] Stabilized weights are computed from the ratio of these probabilities and used to train recurrent encoder and decoder outcome models. The encoder predicts one-step outcomes from \[u^{\mathrm{enc}}_{i,t} = \left[x_{i,t},\;\tilde{y}_{i,t},\;\phi_2(a_{i,t}),\;c_i\right],\] and the decoder performs autoregressive multi-step rollout under the planned future treatment sequence.

14.3 G-Net↩︎

G-Net is a recurrent neural \(g\)-computation baseline [7]. It models the next outcome jointly with the next covariates and then rolls the system forward under a planned treatment sequence. At time \(t\), the input is \[u^{\mathrm{GNet}}_{i,t} = \left[\phi_4(a_{i,t}),\;x_{i,t},\;\tilde{y}_{i,t},\;c_i\right].\] The training objective combines next-outcome prediction with next-covariate prediction: \[\mathcal{L}_{\mathrm{GNet}} = \frac{ \sum_{i,t}m_{i,t} \left(\hat{y}_{i,t+1}-\tilde{y}_{i,t+1}\right)^2 }{ \sum_{i,t}m_{i,t} } + \lambda_x \frac{ \sum_{i,t}m^{x}_{i,t} \left\|\hat{x}_{i,t+1}-x_{i,t+1}\right\|_2^2 }{ \sum_{i,t}m^{x}_{i,t} }.\] At test time, predicted outcomes and covariates are fed back autoregressively, implementing plug-in neural \(g\)-computation under the supplied treatment sequence.

14.4 Counterfactual Recurrent Network↩︎

CRN learns balanced recurrent representations by combining outcome prediction with adversarial treatment prediction [6]. The encoder input is \[u^{\mathrm{CRN}}_{i,t} = \left[\phi_4(a_{i,t-1}),\;x_{i,t},\;\tilde{y}_{i,t},\;c_i\right],\] with a zero previous-treatment vector at \(t=0\). The recurrent representation is mapped to a balanced representation \[b_{i,t}=\operatorname{ELU}(W_b h_{i,t}+c_b).\] The treatment head predicts \(a_{i,t}\) from a gradient-reversed version of \(b_{i,t}\), while the outcome head predicts \(\tilde{y}_{i,t+1}\) from \([b_{i,t},\phi_4(a_{i,t})]\). The objective is \[\mathcal{L}_{\mathrm{CRN}} = \frac{\sum_{i,t}m_{i,t} \left(\hat{y}_{i,t+1}-\tilde{y}_{i,t+1}\right)^2}{\sum_{i,t}m_{i,t}} + \lambda_a\, \mathrm{CE}_{\mathrm{active}}(\hat{a}_{i,t},a_{i,t}).\] A recurrent decoder performs autoregressive rollout under planned treatments.

14.5 Causal Transformer↩︎

Causal Transformer replaces the recurrent backbone of CRN with a multi-input transformer [8]. The treatment, outcome, and covariate streams are initialized as \[u^{a}_{i,t}=W_a\phi_4(a_{i,t-1}),\qquad u^{y}_{i,t}=W_y\tilde{y}_{i,t},\qquad u^{x}_{i,t}=W_x x_{i,t}.\] Transformer blocks apply causal attention within and across streams. The final representation is passed to balanced treatment and outcome heads. The loss is \[\mathcal{L}_{\mathrm{CT}} = \frac{ \sum_{i,t}m_{i,t} \left(\hat{y}_{i,t+1}-\tilde{y}_{i,t+1}\right)^2 }{ \sum_{i,t}m_{i,t} } + \lambda_a \frac{ \sum_{i,t}m_{i,t}\, \mathrm{CE}(\hat{a}_{i,t},a_{i,t}) }{ \sum_{i,t}m_{i,t} }.\] During counterfactual rollout, future covariates are hidden and predicted outcomes are fed back autoregressively.

14.6 G-Transformer↩︎

G-Transformer is a transformer-based neural \(g\)-computation baseline inspired by [9]. It uses treatment, outcome, and covariate streams with transformer attention, but predicts outcomes through a factual \(g\)-computation head rather than an adversarial treatment-balancing head. After the transformer stack, the representation is mapped to \[h^r_{i,t}=\operatorname{ELU}(W_r h_{i,t}+c_r),\] and the one-step head predicts \(\tilde{y}_{i,t+1}\) from \([h^r_{i,t},\phi_4(a_{i,t})]\). The loss is the masked factual MSE, \[\mathcal{L}_{\mathrm{GT}} = \frac{ \sum_{i,t}m_{i,t} \left(\hat{y}_{i,t+1}-\tilde{y}_{i,t+1}\right)^2 }{ \sum_{i,t}m_{i,t} }.\] At test time, the one-step head is applied autoregressively under the planned future treatment sequence.

Table 10: Summary of baseline mechanisms. The baselines cover inverse-probability weighting, neural \(g\)-computation, adversarial balancing, and transformer-based longitudinal sequence modeling.
Method Sequence model Adjustment mechanism Outcome objective
MSM Linear / ridge regressors IPTW via logistic propensities Weighted direct-horizon regression
RMSN LSTM encoder–decoder IPTW via RNN propensities Weighted MSE
G-Net LSTM Neural \(g\)-computation Outcome/covariate MSE
CRN LSTM encoder–decoder Gradient reversal Outcome MSE \(+\) treatment CE
CT Multi-input transformer Gradient reversal Outcome MSE \(+\) treatment CE
GT Multi-input transformer Neural \(g\)-computation Factual outcome MSE

15 Evaluation protocol details↩︎

15.0.0.1 Rolling-origin filtering and indexing.

The raw benchmark generators use trajectory length \(T=60\) and projection horizon \(H=5\). They generate candidate rolling-origin rows using the generator-level minimum-origin setting, currently \(\texttt{min\_t\_obs}=10\) in the benchmark configuration. The shared evaluation layer uses the same minimum observed history length and applies common validity filters: only rows with \[t_{\mathrm{obs}}\ge 10,\qquad t_{\mathrm{obs}}\le 64,\qquad t_{\mathrm{target}}\le 65\] are scored. Thus, \(10\) is the minimum observed history length for both generated candidate origins and reported evaluation rows in the current configuration.

One-step rows use \[t_{\mathrm{obs}}=\texttt{sequence\_lengths}-1, \qquad t_{\mathrm{target}}=\texttt{sequence\_lengths}.\] Horizon-5 rows use the stored rolling-origin index with \[t_{\mathrm{obs}}=\texttt{patient\_current\_t}+1, \qquad t_{\mathrm{target}}=t_{\mathrm{obs}}+5.\] These conventions match the inclusive observation-time convention used in Section 2.1: the query history is observed through \(t_{\mathrm{obs}}\), the first planned treatment is \(a_{t_{\mathrm{obs}}}\), and the target is \(Y_{t_{\mathrm{target}}}\).

15.0.0.2 Support-only normalization and clipping.

Outcome normalization statistics are computed from the support trajectories of each benchmark file. Query targets are never used to estimate normalization statistics. The reported normalized target is \[Y^{\mathrm{eval}}_{i,t} = \operatorname{clip}\!\left( \frac{Y_{i,t}-\mu_Y^{\mathrm{eval}}}{\sigma_Y^{\mathrm{eval}}}, -10,10 \right),\] where \((\mu_Y^{\mathrm{eval}},\sigma_Y^{\mathrm{eval}})\) are computed from the full support set. Reported predictions are expressed in the same evaluation normalization and clipped to \([-20,20]\): \[\widehat Y^{\mathrm{eval}}_{i,t} = \operatorname{clip}\!\left( \frac{\widehat Y^{\mathrm{raw}}_{i,t}-\mu_Y^{\mathrm{eval}}}{\sigma_Y^{\mathrm{eval}}}, -20,20 \right).\] For CausalLongPFN, the model may internally predict in the PFN-context normalization. Before scoring, predictions are converted back through the context outcome scale and then into the shared evaluation normalization.

The reported normalized RMSE is \[\mathrm{RMSE}_{\mathrm{norm}} = \sqrt{ \frac{1}{|\mathcal{I}_{\mathrm{eval}}|} \sum_{(i,t)\in\mathcal{I}_{\mathrm{eval}}} \left( \widehat Y^{\mathrm{eval}}_{i,t} - Y^{\mathrm{eval}}_{i,t} \right)^2 },\] where \(\mathcal{I}_{\mathrm{eval}}\) denotes the set of scored query rows for the corresponding dataset, task, and method.

16 Reproducibility, statistical uncertainty, and compute↩︎

This appendix provides reproducibility, statistical-uncertainty, and compute details for the experiments in Section 3. The model architecture is specified in Appendix 9, the synthetic TSCM prior in Appendix 7, the loss and training procedure in Appendices 1011, the rollout protocol in Appendix 12, the evaluation datasets in Appendix 13, the baseline models in Appendix 14, and the shared evaluation protocol in Appendix 15.

16.0.0.1 Reproducing CausalLongPFN training.

CausalLongPFN is trained entirely on synthetic episodes generated online from the TSCM prior described in Appendix 7. The reported model uses the architecture and optimization hyperparameters in Tables 6 and 7. The implementation, synthetic episode generator, training configuration, rollout code, and evaluation scripts are available at https://github.com/Amirhossein-Zare/causal-long-pfn. No target-domain trajectories are used during CausalLongPFN pretraining.

16.0.0.2 Reproducing benchmark evaluations.

Cancer, warfarin, and HIV are simulated or semi-mechanistic domains with branchable counterfactual labels, as described in Appendix 13. These datasets can be regenerated from the simulator specifications, task-grid configuration, random seeds, and code released at https://github.com/Amirhossein-Zare/causal-long-pfn. For these domains, query labels are obtained by replaying the same patient-specific dynamics under the evaluated treatment sequence. MIMIC-III is a credentialed-access de-identified clinical database and cannot be redistributed with this paper. Reproducing MIMIC-III results therefore requires access through the official MIMIC-III data-use process and the preprocessing protocol described in Appendix 13.4. MIMIC-III evaluation is factual rolling-origin prediction only.

16.0.0.3 Baseline reproducibility.

All baselines are trained using only target-domain support trajectories. Hyperparameters are selected by support-set validation using the search spaces in Table 9. The evaluation runner uses grouped tuning: an initial random search is performed for the first dataset in each \((\text{domain}, n_{\mathrm{sup}})\) group, and the best cached candidate is reused and re-evaluated for later datasets in that group. The selected configuration is then refit on the full support set before query evaluation. Query outcomes are never used for hyperparameter selection. This gives the baselines domain-specific training and validation-based model selection, whereas CausalLongPFN is evaluated frozen without test-time parameter updates, validation-based model selection, or target-domain hyperparameter tuning.

16.0.0.4 Statistical uncertainty.

The aggregation unit is a dataset-level evaluation unit: one normalized RMSE value per method, run, dataset, domain, confounding level \(\gamma\), support size, and prediction task. For long-format prediction outputs, this value is obtained by first aggregating all scored query rows within the unit into a normalized RMSE. For each domain, task, and method, we report the mean normalized RMSE across these evaluation units. When standard errors are reported, they are computed as \[\operatorname{SE}(\hat{m}) = \frac{\operatorname{sd}(m_1,\ldots,m_J)}{\sqrt{J}},\] where \(m_j\) is the normalized RMSE for evaluation unit \(j\) and \(J\) is the number of evaluation units in the aggregation. Domain-balanced summaries are computed by first averaging within each domain and then averaging the resulting domain means equally across domains. For MIMIC-III, these uncertainty summaries describe factual rolling-origin prediction variability and should not be interpreted as uncertainty in counterfactual treatment effects.

16.0.0.5 Compute resources.

The reported CausalLongPFN pretraining configuration uses batch size \(16\) with gradient accumulation over \(16\) micro-batches, giving an effective batch size of \(256\) synthetic episodes. Training is configured for up to \(10000\) optimizer steps, with a session timeout of \(42000\) seconds. Checkpoints are written every \(500\) optimizer steps, and the latest three step checkpoints are retained. The configuration is designed for CUDA-enabled training and supports multi-GPU data parallelism when multiple GPUs are visible. Wall-clock time and the exact number of completed optimizer steps depend on the available hardware and on whether training stops by reaching the optimizer-step budget or the session-timeout limit. Baseline models are trained separately for target tasks and therefore require additional compute proportional to the number of benchmark files, support sizes, and hyperparameter configurations.

17 Data assets, licenses, and ethics↩︎

The proposed CausalLongPFN model, synthetic TSCM prior, and generated synthetic training episodes are new research assets introduced by this work. The paper documents their intended use, causal assumptions, limitations, architecture, training procedure, rollout protocol, and evaluation protocol in Sections 25 and Appendices 716. Code for the model, synthetic data generation, benchmark construction, and evaluation is available at https://github.com/Amirhossein-Zare/causal-long-pfn.

The cancer, HIV, MIMIC-III, and baseline methods are based on previously published benchmarks, simulators, datasets, or modeling frameworks cited in Appendices 13 and 14. This paper credits the original sources used for benchmark construction and baseline comparison. Reused simulator code, preprocessing code, and baseline implementations should be used in accordance with their respective licenses and terms of use.

MIMIC-III is a de-identified credentialed-access clinical database and is not redistributed with this paper. Reproducing MIMIC-III experiments requires obtaining access through the official data-use process and applying the preprocessing protocol described in Appendix 13.4. Results on MIMIC-III are factual rolling-origin prediction results and should not be interpreted as validation of individual counterfactual treatment effects under unobserved ICU interventions.

This work does not involve new human-subject recruitment, prospective interventions, crowdsourcing, or direct interaction with patients. The clinical component uses an existing de-identified database, and all reported MIMIC-III results are aggregate benchmark metrics. CausalLongPFN should be viewed as a research tool for causal sequence modeling and hypothesis generation, not as a standalone clinical decision system.

References↩︎

[1]
D. B. Rubin, “Estimating causal effects of treatments in randomized and nonrandomized studies,” Journal of Educational Psychology, vol. 66, no. 5, pp. 688–701, 1974, doi: 10.1037/h0037350.
[2]
J. M. Robins, “A new approach to causal inference in mortality studies with a sustained exposure period—application to control of the healthy worker survivor effect,” Mathematical Modelling, vol. 7, no. 9–12, pp. 1393–1512, 1986, doi: 10.1016/0270-0255(86)90088-6.
[3]
J. M. Robins, M. A. Hernán, and B. Brumback, “Marginal structural models and causal inference in epidemiology,” Epidemiology, vol. 11, no. 5, pp. 550–560, 2000, doi: 10.1097/00001648-200009000-00011.
[4]
M. A. Hernán and J. M. Robins, Causal inference: What if. Boca Raton: Chapman & Hall/CRC, 2020.
[5]
B. Lim, “Forecasting treatment responses over time using recurrent marginal structural networks,” in Advances in neural information processing systems, 2018, vol. 31, [Online]. Available: https://proceedings.neurips.cc/paper_files/paper/2018/file/56e6a93212e4482d99c84a639d254b67-Paper.pdf.
[6]
I. Bica, A. M. Alaa, J. Jordon, and M. van der Schaar, “Estimating counterfactual treatment outcomes over time through adversarially balanced representations,” in International conference on learning representations, 2020, [Online]. Available: https://openreview.net/forum?id=BJg866NFvB.
[7]
R. Li et al., G-Net: A recurrent network approach to G-computation for counterfactual prediction under a dynamic treatment regime,” in Proceedings of machine learning for health, 2021, vol. 158, pp. 282–299, [Online]. Available: https://proceedings.mlr.press/v158/li21a.html.
[8]
V. Melnychuk, D. Frauen, and S. Feuerriegel, “Causal transformer for estimating counterfactual outcomes,” in Proceedings of the 39th international conference on machine learning, 2022, vol. 162, pp. 15293–15329, [Online]. Available: https://proceedings.mlr.press/v162/melnychuk22a.html.
[9]
H. Xiong, F. Wu, L. Deng, M. Su, and L. H. Lehman, G-Transformer: Counterfactual outcome prediction under dynamic and time-varying treatment regimes,” in Proceedings of the 9th machine learning for healthcare conference, 2024, vol. 252, [Online]. Available: https://proceedings.mlr.press/v252/xiong24a.html.
[10]
S. Müller, N. Hollmann, S. Pineda Arango, J. Grabocka, and F. Hutter, “Transformers can do Bayesian inference,” in International conference on learning representations, 2022, [Online]. Available: https://openreview.net/forum?id=KSugKcbNf9.
[11]
N. Hollmann, S. Müller, K. Eggensperger, and F. Hutter, TabPFN: A transformer that solves small tabular classification problems in a second,” in International conference on learning representations, 2023, [Online]. Available: https://openreview.net/forum?id=cp5PvcI6w8_.
[12]
S. Dooley, G. S. Khurana, C. Mohapatra, S. Naidu, and C. White, ForecastPFN: Synthetically-trained zero-shot forecasting.” 2023, [Online]. Available: https://arxiv.org/abs/2311.01933.
[13]
E. O. Taga, M. E. Ildiz, and S. Oymak, TimePFN: Effective multivariate time series forecasting with synthetic data,” in Proceedings of the AAAI conference on artificial intelligence, 2025, vol. 39, pp. 20761–20769.
[14]
J. Robertson, A. Reuter, S. Guo, N. Hollmann, F. Hutter, and B. Schölkopf, Do-PFN: In-context learning for causal effect estimation.” 2025, [Online]. Available: https://arxiv.org/abs/2506.06039.
[15]
V. Balazadeh et al., “CausalPFN: Amortized causal effect estimation via in-context learning,” in Advances in neural information processing systems, 2025, vol. 38.
[16]
Y. Ma, D. Frauen, E. Javurek, and S. Feuerriegel, “Foundation models for causal inference via prior-data fitted networks,” arXiv preprint arXiv:2506.10914, 2025, [Online]. Available: https://arxiv.org/abs/2506.10914.
[17]
D. Thumm and Y. Chen, “Interventional time series priors for causal foundation models,” in 1st ICLR workshop on time series in the age of large models, 2026, [Online]. Available: https://openreview.net/forum?id=JbTgx2L9Z2.
[18]
J. Pearl, Causality: Models, reasoning, and inference, 2nd ed. Cambridge: Cambridge University Press, 2009.
[19]
J. Peters, D. Janzing, and B. Schölkopf, Elements of causal inference: Foundations and learning algorithms. Cambridge, MA: The MIT Press, 2017.
[20]
T. Nagler, “Statistical foundations of prior-data fitted networks,” in Proceedings of the 40th international conference on machine learning, 2023, vol. 202, pp. 25660–25676, [Online]. Available: https://proceedings.mlr.press/v202/nagler23a.html.
[21]
A. Vaswani et al., “Attention is all you need,” in Advances in neural information processing systems, 2017, vol. 30, [Online]. Available: https://proceedings.neurips.cc/paper_files/paper/2017/file/3f5ee243547dee91fbd053c1c4a845aa-Paper.pdf.
[22]
C. M. Bishop, Copyright © 1994, Christopher M. Bishop. This work is licensed under a Creative Commons Attribution-NonCommercial-NoDerivatives 4.0 International License (https://creativecommons.org/licenses/by-nc-nd/4.0/).Mixture density networks,” Aston University, Birmingham, 1994.
[23]
C. Geng, H. Paganetti, and C. Grassberger, “Prediction of treatment response for combined chemo- and radiation therapy for non-small cell lung cancer patients using a bio-mathematical model,” Scientific Reports, vol. 7, p. 13542, 2017, doi: 10.1038/s41598-017-13646-z.
[24]
N. H. G. Holford, “Clinical pharmacokinetics and pharmacodynamics of warfarin: Understanding the dose-effect relationship,” Clinical Pharmacokinetics, vol. 11, no. 6, pp. 483–504, 1986, doi: 10.2165/00003088-198611060-00005.
[25]
A.-K. Hamberg et al., “A pharmacometric model describing the relationship between warfarin dose and INR response with respect to variations in CYP2C9, VKORC1, and age,” Clinical Pharmacology & Therapeutics, vol. 87, no. 6, pp. 727–734, 2010, doi: https://doi.org/10.1038/clpt.2010.37.
[26]
International Warfarin Pharmacogenetics Consortium, “Estimation of the warfarin dose with clinical and pharmacogenetic data,” New England Journal of Medicine, vol. 360, no. 8, pp. 753–764, 2009, doi: 10.1056/NEJMoa0809329.
[27]
B. M. Adams, H. T. Banks, H.-D. Kwon, and H. T. Tran, “Dynamic multidrug therapies for HIV: Optimal and STI control approaches,” Mathematical Biosciences and Engineering, vol. 1, no. 2, pp. 223–241, 2004, doi: 10.3934/mbe.2004.1.223.
[28]
J. Miller et al., WhyNot.” Software package; Zenodo, 2020, doi: 10.5281/zenodo.3875775.
[29]
A. Johnson, T. Pollard, and R. Mark, Version 1.4MIMIC-III Clinical Database,” PhysioNet, Sep. 2016, doi: 10.13026/C2XW26.
[30]
A. E. W. Johnson et al., MIMIC-III, a freely accessible critical care database,” Scientific Data, vol. 3, p. 160035, 2016, doi: 10.1038/sdata.2016.35.
[31]
A. L. Goldberger, L. A. N. Amaral, L. Glass, et al., “PhysioBank, PhysioToolkit, and PhysioNet: Components of a new research resource for complex physiologic signals,” Circulation, vol. 101, no. 23, pp. e215–e220, 2000.
[32]
S. Wang, M. B. A. McDermott, G. Chauhan, M. C. Hughes, T. Naumann, and M. Ghassemi, MIMIC-Extract: A data extraction, preprocessing, and representation pipeline for MIMIC-III,” in Proceedings of the ACM conference on health, inference, and learning, 2020, pp. 222–235, doi: 10.1145/3368555.3384469.
[33]
H. Harutyunyan, H. Khachatrian, D. C. Kale, G. Ver Steeg, and A. Galstyan, “Multitask learning and benchmarking with clinical time series data,” Scientific Data, vol. 6, p. 96, 2019, doi: 10.1038/s41597-019-0103-9.
[34]
T. Gneiting, F. Balabdaoui, and A. E. Raftery, “Probabilistic forecasts, calibration and sharpness,” Journal of the Royal Statistical Society: Series B (Statistical Methodology), vol. 69, no. 2, pp. 243–268, 2007, doi: 10.1111/j.1467-9868.2007.00587.x.
[35]
T. Gneiting and A. E. Raftery, “Strictly proper scoring rules, prediction, and estimation,” Journal of the American Statistical Association, vol. 102, no. 477, pp. 359–378, 2007, doi: 10.1198/016214506000001437.
[36]
I. Loshchilov and F. Hutter, “Decoupled weight decay regularization,” in International conference on learning representations, 2019, [Online]. Available: https://openreview.net/forum?id=Bkg6RiCqY7.