Skip to content

Commit 1e23cf3

Browse files
committed
Removed scenario from CausalDAG
1 parent 977aa67 commit 1e23cf3

5 files changed

Lines changed: 25 additions & 19 deletions

File tree

causal_testing/main.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -313,7 +313,10 @@ def create_causal_test(self, test: dict, base_test: BaseTestCase) -> CausalTestC
313313
base_test_case=base_test,
314314
treatment_value=test.get("treatment_value"),
315315
control_value=test.get("control_value"),
316-
adjustment_set=test.get("adjustment_set", self.causal_specification.causal_dag.identification(base_test)),
316+
adjustment_set=test.get(
317+
"adjustment_set",
318+
self.causal_specification.causal_dag.identification(base_test, self.scenario.hidden_variables()),
319+
),
317320
df=filtered_df,
318321
effect_modifiers=None,
319322
formula=test.get("formula"),

causal_testing/specification/causal_dag.py

Lines changed: 6 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,7 @@
1010

1111
from causal_testing.testing.base_test_case import BaseTestCase
1212

13-
from .scenario import Scenario
14-
from .variable import Output
13+
from .variable import Variable
1514

1615
Node = Union[str, int] # Node type hint: A node is a string or an int
1716

@@ -489,16 +488,7 @@ def get_backdoor_graph(self, treatments: list[str]) -> CausalDAG:
489488
backdoor_graph.add_edges_from(filter(lambda x: x not in outgoing_edges, self.edges))
490489
return backdoor_graph
491490

492-
@staticmethod
493-
def remove_hidden_adjustment_sets(minimal_adjustment_sets: list[str], scenario: Scenario):
494-
"""Remove variables labelled as hidden from adjustment set(s)
495-
496-
:param minimal_adjustment_sets: list of minimal adjustment set(s) to have hidden variables removed from
497-
:param scenario: The modelling scenario which informs the variables that are hidden
498-
"""
499-
return [adj for adj in minimal_adjustment_sets if all(not scenario.variables.get(x).hidden for x in adj)]
500-
501-
def identification(self, base_test_case: BaseTestCase, scenario: Scenario = None):
491+
def identification(self, base_test_case: BaseTestCase, avoid_variables: set[Variable] = None):
502492
"""Identify and return the minimum adjustment set
503493
504494
:param base_test_case: A base test case instance containing the outcome_variable and the
@@ -523,8 +513,10 @@ def identification(self, base_test_case: BaseTestCase, scenario: Scenario = None
523513
else:
524514
raise ValueError("Causal effect should be 'total' or 'direct'")
525515

526-
if scenario is not None:
527-
minimal_adjustment_sets = self.remove_hidden_adjustment_sets(minimal_adjustment_sets, scenario)
516+
if avoid_variables is not None:
517+
minimal_adjustment_sets = [
518+
adj for adj in minimal_adjustment_sets if not {x.name for x in avoid_variables}.intersection(adj)
519+
]
528520

529521
minimal_adjustment_set = min(minimal_adjustment_sets, key=len, default=set())
530522
return set(minimal_adjustment_set)

causal_testing/specification/scenario.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,3 +31,11 @@ def __init__(self, variables: Iterable[Variable], constraints: set[str] = None):
3131
self.constraints = set(constraints)
3232
else:
3333
self.constraints = set()
34+
35+
def hidden_variables(self) -> set[Variable]:
36+
"""Get the set of hidden variables
37+
38+
:return The variables marked as hidden.
39+
:rtype: {Variable}
40+
"""
41+
return {v for v in self.variables.values() if v.hidden}

causal_testing/surrogate/causal_surrogate_assisted.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -133,7 +133,9 @@ def generate_surrogates(
133133
to_var = specification.scenario.variables.get(v)
134134
base_test_case = BaseTestCase(from_var, to_var)
135135

136-
minimal_adjustment_set = specification.causal_dag.identification(base_test_case, specification.scenario)
136+
minimal_adjustment_set = specification.causal_dag.identification(
137+
base_test_case, specification.scenario.hidden_variables()
138+
)
137139

138140
surrogate = CubicSplineRegressionEstimator(
139141
base_test_case,

tests/specification_tests/test_causal_dag.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import unittest
22
import os
3-
import shutil, tempfile
3+
import shutil
4+
import tempfile
45
import networkx as nx
56
from causal_testing.specification.causal_dag import CausalDAG, close_separator, list_all_min_sep
67
from causal_testing.specification.scenario import Scenario
@@ -420,10 +421,10 @@ def test_hidden_varaible_adjustment_sets(self):
420421
m = Input("M", int)
421422

422423
scenario = Scenario(variables={z, x, m})
423-
adjustment_sets = causal_dag.identification(BaseTestCase(x, m), scenario)
424+
adjustment_sets = causal_dag.identification(BaseTestCase(x, m), scenario.hidden_variables())
424425

425426
z.hidden = True
426-
adjustment_sets_with_hidden = causal_dag.identification(BaseTestCase(x, m), scenario)
427+
adjustment_sets_with_hidden = causal_dag.identification(BaseTestCase(x, m), scenario.hidden_variables())
427428

428429
self.assertNotEqual(adjustment_sets, adjustment_sets_with_hidden)
429430

0 commit comments

Comments
 (0)