A probabilistic program is also a causal model: each sample site is a variable computed from the ones before it. That lets us ask three kinds of question (Pearl's ladder): what we expect having seen something, what would happen if we did something, and what would have happened had things been different.
(ns foerster.counterfactuals
(:require [org.replikativ.foerster.core :as infer]
[org.replikativ.foerster.counterfactual :as cf]
[org.replikativ.foerster.dist :as dist]
[org.replikativ.foerster.effects :refer [sample]]
[org.replikativ.foerster.mechanism :as mech]
[org.replikativ.foerster.measure :as m]
[org.replikativ.foerster.random :as random]
[org.replikativ.foerster.trace :as trace]
[org.replikativ.spindel.core :as sp]
[org.replikativ.spindel.select :as select]
[org.replikativ.spindel.spin.cps :refer [spin]]))(def world (sp/create-execution-context))(random/set-seed! 3)nil
A hidden common cause Z drives both X and Y, and X also affects Y: Z ~ N(0, 1), X ~ N(Z, 1), Y ~ N(X + Z, 1).
(defn scm []
(spin
(let [z (sample (dist/normal 0.0 1.0) :id :z)
x (sample (dist/normal z 1.0) :id :x)
y (sample (dist/normal (+ x z) 1.0) :id :y)]
{:z z :x x :y y})))(defn expected-y [policy]
(:mean (infer/query (sp/with-context world
@(infer/importance-sampling (scm) 4000 {:policy (trace/policy policy)}))
:y)))Seeing X = 1 conditions on it: it says something about Z too (Z is likely positive), so E[Y | X = 1] = 1.5. A constraint fixes a site to a value and scores it, like an observation:
(expected-y {:constraints {:x 1.0}})1.4802572814461274
Doing X = 1 sets it, cutting it off from Z: E[Y | do(X = 1)] = 1. An intervention replaces the site's mechanism and scores nothing:
(expected-y {:interventions {:x {:do 1.0}}})1.0012518197243876
Interventions are keyed by address or by spindel selector, so one intervention can reach every site of a kind, or every site under a path:
(expected-y {:interventions {(select/id :x) {:do 1.0}}})0.9689145696388346
Besides :do, an intervention can set a new mechanism (:dist), shift the old one (:shift), or compute the value from the choices made so far (:policy, a function of them):
{:new-mechanism (expected-y {:interventions {:x {:dist (dist/normal 3.0 0.1)}}})
:shifted (expected-y {:interventions {:x {:shift 2.0}}})
:policy (expected-y {:interventions {:x {:policy (fn [choices] (* 2.0 (:z choices)))}}})}{:new-mechanism 3.0047276535061123,
:shifted 1.9737024667438372,
:policy 0.043367976340663764}
(Exact: 3, 2 and 0.)
A counterfactual asks about this situation, not a new one: we saw X = 1 and Y = 2.5 — had X been 2, what would Y have been? The answer keeps everything about the situation we saw — including the noise that made Y what it was — and changes only X.
To do that, each site is read as a structural equation x = f(u) of exogenous noise u (foerster.mechanism). counterfactual:
(def pairs
(let [p (promise)]
((cf/counterfactual (sp/with-context world (scm))
{:evidence {:x 1.0 :y 2.5}
:interventions {:x {:do 2.0}}
:particles 50})
#(deliver p %) #(deliver p %))
(deref p 60000 :timeout)))counterfactual takes the model spin, as gfi does — it runs it in a factual and a counterfactual world per particle — and returns a CPS operation resolving factual/counterfactual pairs, one per particle, with weights:
(select-keys (first pairs) [:factual :counterfactual :weight]){:factual {:z -0.46724440593680316, :x 1.0, :y 2.5},
:counterfactual {:z -0.46724440593680316, :x 2.0, :y 3.5},
:weight -4.84930541613059}
Z is uncertain given the evidence, but Y − X = Z + noise is known exactly (1.5), so every particle agrees: had X been 2, Y would have been 3.5.
(set (map #(/ (Math/round (* 1e9 (:y (:counterfactual %)))) 1e9) pairs))#{3.5}
X ~ flip(0.5), and Y = true for sure when X is, with probability 0.3 otherwise. We saw X and Y both true. The probability of necessity — would Y have been false had X been false? — is 0.7.
(defn binary []
(sp/with-context world
(spin
(let [x (sample (dist/flip 0.5) :id :x)
y (sample (dist/flip (if x 1.0 0.3)) :id :y)]
{:x x :y y}))))(let [p (promise)]
((cf/counterfactual (binary) {:evidence {:x true :y true}
:interventions {:x {:do false}}
:particles 1500})
#(deliver p %) #(deliver p %))
(let [pairs (deref p 60000 [])
ws (m/normalize-log-weights (mapv :weight pairs))]
(reduce + (map (fn [pair w] (if (:y (:counterfactual pair)) 0.0 w)) pairs ws))))0.7120000000000067
A discrete site has many noise values for one outcome; abduction draws one from its posterior. Which convention a discrete mechanism uses (here the inverse CDF, x = F⁻¹(u)) changes counterfactual answers and cannot be told from data, so foerster fixes and documents it (foerster.mechanism) rather than guessing.
A site that exists only in the counterfactual world — reached through a branch the factual run did not take — has no factual noise; it is drawn fresh and reported in :unaligned.