CaLM: Causal Language Modeling for Counterfactual World Simulations
Abstract
Counterfactual trajectory simulation asks how a sequence would have unfolded had selected actions been different. Autoregressive Transformers provide a powerful framework for sequence modeling by learning the conditional distribution of each token given the preceding tokens. However, fitting these conditional distributions does not ensure that simulations respect causal relationships or correctly account for unobserved confounding when evaluating interventions. We formalize counterfactual trajectory simulation in the language of structural causal models. To evaluate how reliably a simulator class answers counterfactual queries, we introduce *counterfactual coverage*, which holds when the true counterfactual value lies between the lower and upper bounds of the class's estimates. We show theoretically that next-token prediction alone does not guarantee counterfactual coverage for vanilla Transformers. We develop CaLM (Causal Language Modeling), a Transformer-based simulator that incorporates constraints specified by an arbitrary causal graph into its attention mechanism. We establish counterfactual coverage for the CaLM simulator class and demonstrate improved counterfactual prediction across chess, buyer-seller negotiation, autonomous driving, clinical trajectories, and poker.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.