π Full Documentation
π» GitHub
π Paper
PyTorch Implementation on Paper Foundation Models for Causal Inference via Prior-Data Fitted Networks
In this paper, we introduce CausalFM, a comprehensive framework for training PFN-based foundation models in various causal inference settings.
CausalFM provides a unified framework for training foundation models across multiple causal inference tasks, including:
- Standard CATE estimation setting
- Instrumental Variables (IV) setting
- Front-door adjustment setting
This repository contains dataset generation pipelines, model implementations, and training/evaluation scripts.
Clone the repository and install dependencies:
git clone https://github.com/yccm/CausalFM-toolkit.git
cd CausalFM-toolkit
conda create -n causalfm python=3.10
conda activate causalfm
pip install -r requirements.txtCausalFM can be used as a library with a clean, intuitive API:
import causalfm
# Load a pretrained model
model = causalfm.StandardCATEModel.from_pretrained("checkpoints/best_model.pth")
# Estimate CATE for new samples
result = model.estimate_cate(x_train, a_train, y_train, x_test)
cate_estimates = result['cate']Generate synthetic datasets for training and evaluation:
from causalfm.data import StandardCATEGenerator, IVDataGenerator, FrontdoorDataGenerator
# Standard CATE data
generator = StandardCATEGenerator(num_samples=1024, num_features=10, seed=42)
df = generator.generate()
# Generate multiple datasets
generator.generate_multiple(num_datasets=10, output_dir="data/standard/")
# Instrumental Variables data
iv_generator = IVDataGenerator(
num_samples=1024,
num_features=10,
instrument_type='binary', # or 'continuous'
seed=42
)
iv_df = iv_generator.generate()
# Front-door adjustment data
fd_generator = FrontdoorDataGenerator(
num_samples=1024,
num_features=10,
num_confounders=5,
seed=42
)
fd_df = fd_generator.generate()Train models using the Trainer classes:
from causalfm.training import StandardCATETrainer, TrainingConfig
# Using configuration object
config = TrainingConfig(
data_path="data/standard/*.csv",
epochs=100,
batch_size=16,
learning_rate=0.001,
save_dir="checkpoints/standard"
)
trainer = StandardCATETrainer(config)
trainer.train()
# Or use simplified interface
trainer = StandardCATETrainer.from_args(
data_path="data/standard/*.csv",
epochs=100,
batch_size=16,
save_dir="checkpoints/standard"
)
trainer.train()Load pretrained models and run inference:
from causalfm.models import StandardCATEModel, IVModel, FrontdoorModel
import torch
# Standard CATE Model
model = StandardCATEModel.from_pretrained("checkpoints/best_model.pth")
# Prepare data
x_train = torch.randn(800, 10) # Training covariates
a_train = torch.randint(0, 2, (800,)).float() # Training treatments
y_train = torch.randn(800) # Training outcomes
x_test = torch.randn(200, 10) # Test covariates
# Estimate CATE
result = model.estimate_cate(x_train, a_train, y_train, x_test)
cate = result['cate'] # Shape: (200,)
# Access GMM distribution parameters (for uncertainty)
pi = result['gmm_pi'] # Mixture weights
mu = result['gmm_mu'] # Means
sigma = result['gmm_sigma'] # Standard deviations
# IV Model
iv_model = IVModel.from_pretrained("checkpoints/iv_model.pth")
result = iv_model.estimate_cate(x_train, z_train, a_train, y_train, x_test)
# Front-door Model
fd_model = FrontdoorModel.from_pretrained("checkpoints/fd_model.pth")
result = fd_model.estimate_cate(x_train, m_train, a_train, y_train, x_test)Evaluate models using standard metrics:
from causalfm.evaluation import compute_pehe, compute_ate_error
from causalfm.data import normalize_data
import pandas as pd
import torch
# Load test data
df = pd.read_csv("data/test/test_dataset_1.csv")
x_cols = [c for c in df.columns if c.startswith('x')]
# Normalize data (important for consistency with training!)
X_norm, Y_norm, x_scaler, y_scaler = normalize_data(
df[x_cols].values,
df['outcome'].values,
df['y0'].values,
df['y1'].values
)
# Prepare tensors
X = torch.FloatTensor(X_norm)
A = torch.FloatTensor(df['treatment'].values).unsqueeze(1)
Y = torch.FloatTensor(Y_norm).unsqueeze(1)
# Get normalized ITE for evaluation
from causalfm.data import normalize_ite
true_ite_norm, _ = normalize_ite(df['y0'].values, df['y1'].values, y_scaler)
# Split and evaluate
n_train = int(0.8 * len(X))
model = StandardCATEModel.from_pretrained("checkpoints/best_model.pth")
result = model.estimate_cate(X[:n_train], A[:n_train], Y[:n_train], X[n_train:])
# Compute metrics
pehe = compute_pehe(result['cate'].cpu().numpy(), true_ite_norm[n_train:])
print(f"PEHE: {pehe:.4f}")For backward compatibility, you can also use the original script-based approach:
Standard CATE:
cd DATA_standard
python gen_standard_syn.py Instrumental Variables (IV):
cd DATA_IV
python gen_iv_data_binary.py # Binary Instrument
python gen_iv_data_conti.py # Continuous InstrumentFront-door adjustment:
cd DATA_FD
python gen_frontdoor.pyStandard CATE:
python src/tabpfn/train_standard/training_standard.py Instrumental Variables (IV):
python src/tabpfn/train_iv/training_iv_binary.py
python src/tabpfn/train_iv/training_iv_conti.pyFront-door adjustment:
python src/tabpfn/train_fd/training_fd.pyβββ evaluation/notebook/
β βββ test_fd.ipynb # Front-door evaluation
β βββ test_iv_binary.ipynb # Binary IV evaluation
β βββ test_iv_conti.ipynb # Continuous IV evaluation
β βββ test_jobs.ipynb # Jobs dataset evaluation
β βββ test_standard_cate.ipynb # Standard CATE evaluation
CausalFM-toolkit/
βββ causalfm/ # Main package (new library interface)
β βββ __init__.py
β βββ data/ # Data generation and loading
β β βββ generators/ # Dataset generators
β β βββ loaders/ # PyTorch data loaders
β βββ models/ # Model wrappers
β β βββ standard.py # StandardCATEModel
β β βββ iv.py # IVModel
β β βββ frontdoor.py # FrontdoorModel
β βββ training/ # Training utilities
β β βββ base.py # BaseTrainer
β β βββ standard.py # StandardCATETrainer
β β βββ iv.py # IVTrainer
β β βββ frontdoor.py # FrontdoorTrainer
β βββ evaluation/ # Evaluation metrics
β βββ metrics.py # PEHE, ATE error, etc.
βββ src/tabpfn/ # Core TabPFN-based models
β βββ model/
β βββ causalFM.py # Standard CATE model
β βββ causalFM4IV.py # IV model
β βββ causalFM4FD.py # Front-door model
βββ DATA_standard/ # Standard CATE data
βββ DATA_IV/ # IV data
βββ DATA_FD/ # Front-door data
βββ evaluation/notebook/ # Evaluation notebooks
For comprehensive guides, tutorials, and API reference, visit our documentation:
π https://causalfm-toolkit.readthedocs.io
The documentation includes:
- Installation Guide - Detailed setup instructions
- Quick Start - Get started in 5 minutes
- Tutorials - Step-by-step learning path
- User Guides - In-depth coverage of all features
- API Reference - Complete API documentation
- Examples - Complete working examples
If you find this repository useful, please cite our paper:
@article{ma2025foundation,
title={Foundation Models for Causal Inference via Prior-Data Fitted Networks},
author={Ma, Yuchen and Frauen, Dennis and Javurek, Emil and Feuerriegel, Stefan},
journal={arXiv preprint arXiv:2506.10914},
year={2025}
}This repo is based on the implementation of TabPFN