Skip to content

Commit 3b5b62a

Browse files
authored
Merge pull request #389 from CITCOM-project/jmafoster1/hotfix
Fixed numpy linalg error
2 parents e6a19c7 + 8f04b99 commit 3b5b62a

2 files changed

Lines changed: 109 additions & 10 deletions

File tree

causal_testing/discovery/abstract_discovery.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from itertools import permutations
1111

1212
import networkx as nx
13+
import numpy as np
1314
import pandas as pd
1415
import rustworkx as rx
1516

@@ -207,28 +208,27 @@ def evaluate_tests(self, causal_dag: CausalDAG) -> pd.DataFrame:
207208

208209
results = []
209210

210-
# We use "silent=False" here to allow for inestimable edges, but it'd be good to have a more stringent
211-
# error catching strategy to catch "genuine" problems (e.g. to do with the structure of the data)
212-
for test_case, result in zip(ctf.test_cases, ctf.run_tests(silent=False)):
213-
if result.effect_estimate is None:
211+
for test_case, result in zip(ctf.test_cases, ctf.test_cases):
212+
try:
213+
result = test_case.execute_test()
214214
results.append(
215215
{
216-
"result": TestResult.INESTIMABLE,
216+
"result": (
217+
TestResult.PASS if test_case.expected_causal_effect.apply(result) else TestResult.FAIL
218+
),
217219
"expected_effect": test_case.expected_causal_effect.__class__.__name__,
218220
"treatment": test_case.base_test_case.treatment_variable.name,
219221
"outcome": test_case.base_test_case.outcome_variable.name,
222+
"effect": effect_direction(result),
220223
}
221224
)
222-
else:
225+
except np.linalg.LinAlgError:
223226
results.append(
224227
{
225-
"result": (
226-
TestResult.PASS if test_case.expected_causal_effect.apply(result) else TestResult.FAIL
227-
),
228+
"result": TestResult.INESTIMABLE,
228229
"expected_effect": test_case.expected_causal_effect.__class__.__name__,
229230
"treatment": test_case.base_test_case.treatment_variable.name,
230231
"outcome": test_case.base_test_case.outcome_variable.name,
231-
"effect": effect_direction(result),
232232
}
233233
)
234234

tests/discovery_tests/test_abstract_discovery.py

Lines changed: 99 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -176,6 +176,105 @@ def test_write_dot_invalid_independence_outcome(self):
176176
with self.assertRaises(ValueError):
177177
abstract_discovery.write_dot(dag, "dag.dot")
178178

179+
def test_evaluate_tests_invalid_datatype(self):
180+
scarf_df = pd.read_csv("tests/resources/data/scarf_data.csv")
181+
scarf_df["completed"] = pd.to_datetime(["2026-01-01" for _ in range(len(scarf_df))], format="%Y-%m-%d")
182+
183+
dag = CausalDAG()
184+
dag.add_nodes_from(scarf_df.columns)
185+
dag.add_edges_from([("length_in", "completed"), ("large_gauge", "completed")])
186+
187+
abstract_discovery = AbstractDiscovery(scarf_df)
188+
with self.assertRaises(ValueError):
189+
abstract_discovery.evaluate_tests(dag)
190+
191+
def test_evaluate_tests_inestimable(self):
192+
scarf_df = pd.read_csv("tests/resources/data/scarf_data.csv")
193+
scarf_df["completed"] = scarf_df["completed"].astype(bool)
194+
scarf_df = scarf_df.loc[(scarf_df["completed"]) & (scarf_df["color"] != "grey")]
195+
196+
dag = CausalDAG()
197+
dag.add_nodes_from(scarf_df.columns)
198+
dag.add_edges_from([("length_in", "completed"), ("large_gauge", "completed")])
199+
200+
abstract_discovery = AbstractDiscovery(scarf_df)
201+
test_results = abstract_discovery.evaluate_tests(dag)
202+
expected_results = pd.DataFrame(
203+
[
204+
{
205+
"result": TestResult.PASS,
206+
"expected_effect": "NoEffect",
207+
"treatment": "length_in",
208+
"outcome": "large_gauge",
209+
"effect": "negative",
210+
},
211+
{
212+
"result": TestResult.PASS,
213+
"expected_effect": "NoEffect",
214+
"treatment": "large_gauge",
215+
"outcome": "length_in",
216+
"effect": "negative",
217+
},
218+
{
219+
"result": TestResult.PASS,
220+
"expected_effect": "NoEffect",
221+
"treatment": "length_in",
222+
"outcome": "color",
223+
"effect": None,
224+
},
225+
{
226+
"result": TestResult.PASS,
227+
"expected_effect": "NoEffect",
228+
"treatment": "color",
229+
"outcome": "length_in",
230+
"effect": None,
231+
},
232+
{
233+
"result": TestResult.FAIL,
234+
"expected_effect": "SomeEffect",
235+
"treatment": "length_in",
236+
"outcome": "completed",
237+
"effect": "positive",
238+
},
239+
{
240+
"result": TestResult.PASS,
241+
"expected_effect": "NoEffect",
242+
"treatment": "large_gauge",
243+
"outcome": "color",
244+
"effect": None,
245+
},
246+
{
247+
"result": TestResult.PASS,
248+
"expected_effect": "NoEffect",
249+
"treatment": "color",
250+
"outcome": "large_gauge",
251+
"effect": None,
252+
},
253+
{
254+
"result": TestResult.FAIL,
255+
"expected_effect": "SomeEffect",
256+
"treatment": "large_gauge",
257+
"outcome": "completed",
258+
"effect": "negative",
259+
},
260+
{
261+
"result": TestResult.INESTIMABLE,
262+
"expected_effect": "NoEffect",
263+
"treatment": "color",
264+
"outcome": "completed",
265+
"effect": None,
266+
},
267+
{
268+
"result": TestResult.INESTIMABLE,
269+
"expected_effect": "NoEffect",
270+
"treatment": "completed",
271+
"outcome": "color",
272+
"effect": None,
273+
},
274+
]
275+
)
276+
pd.testing.assert_frame_equal(test_results, expected_results)
277+
179278
def test_evaluate_tests(self):
180279
scarf_df = pd.read_csv("tests/resources/data/scarf_data.csv")
181280

0 commit comments

Comments
 (0)