ARotting's picture
Publish Cross-fitted doubly robust causal-effect estimators
6b67395 verified
Raw
History Blame Contribute Delete
2.46 kB
from __future__ import annotations
import gradio as gr
import plotly.graph_objects as go
from causal import estimate_effects, generate_scm
def simulate(
samples: int, confounding: float, seed: int
) -> tuple[go.Figure, dict]:
sample = generate_scm(
samples=int(samples),
seed=int(seed),
confounding=float(confounding),
)
results = estimate_effects(
sample,
propensity_correct=True,
outcome_correct=True,
folds=5,
seed=int(seed) + 1,
)
labels = ["Ground truth", "Naive association", "IPW", "Outcome model", "AIPW"]
values = [
results["truth"],
results["naive"],
results["ipw"],
results["outcome_regression"],
results["aipw"],
]
colors = ["#22c55e", "#ef4444", "#38bdf8", "#a78bfa", "#f59e0b"]
figure = go.Figure(
go.Bar(
x=labels,
y=values,
marker_color=colors,
text=[f"{value:.3f}" for value in values],
textposition="outside",
)
)
figure.update_layout(
title="Association versus causal effect",
yaxis_title="Estimated average treatment effect",
template="plotly_dark",
margin=dict(t=65, b=40),
)
return figure, {
"true_ate": round(results["truth"], 4),
"naive_bias": round(results["naive"] - results["truth"], 4),
"aipw_bias": round(results["aipw"] - results["truth"], 4),
"aipw_95_percent_interval": [
round(results["aipw_ci_low"], 4),
round(results["aipw_ci_high"], 4),
],
}
with gr.Blocks(title="Causal Forge") as demo:
gr.Markdown(
"# Causal Forge\n"
"Create a confounded observational study with known counterfactual truth, "
"then watch causal estimators try to recover the real treatment effect."
)
with gr.Row():
samples = gr.Slider(1_000, 20_000, 5_000, step=500, label="Study size")
confounding = gr.Slider(
0.0, 1.8, 1.0, step=0.1, label="Confounding strength"
)
seed = gr.Number(2043, precision=0, label="Random seed")
run = gr.Button("Run observational study", variant="primary")
chart = gr.Plot()
metrics = gr.JSON()
run.click(simulate, [samples, confounding, seed], [chart, metrics])
demo.load(simulate, [samples, confounding, seed], [chart, metrics])
if __name__ == "__main__":
demo.launch()