From b19edd7b6a00a6a968c010dcfac6661f421f8c4a Mon Sep 17 00:00:00 2001 From: Jake Newman Date: Tue, 14 Jul 2026 11:13:02 +0100 Subject: [PATCH 1/3] Context node or concatenation --- causal_testing/__main__.py | 22 ++++++++++++++----- .../discovery/abstract_discovery.py | 3 +++ causal_testing/main.py | 8 +++++++ 3 files changed, 28 insertions(+), 5 deletions(-) diff --git a/causal_testing/__main__.py b/causal_testing/__main__.py index 3ecd443a..d6929e46 100644 --- a/causal_testing/__main__.py +++ b/causal_testing/__main__.py @@ -58,11 +58,23 @@ def main() -> None: kwargs[split[0]] = split[1] logging.info("Discovering causal structure") - # Need to reset index to allow for multiple files having the same index (i.e. starting at zero). - # Otherwise you end up with duplicate indices, which causes problems further down the line - df = pd.concat([pd.read_csv(path) for path in args.data_paths]).reset_index() - if args.variables: - df = df[args.variables] + + if args.context: + dfs = [] + for i, path in enumerate(args.data_paths): + temp_df = pd.read_csv(path) + temp_df['file_index'] = i + dfs.append(temp_df) + df = pd.concat(dfs, ignore_index=True) + + if args.variables: + df = df[list(set(args.variables + ['file_index']))] + else: + # Need to reset index to allow for multiple files having the same index (i.e. starting at zero). + # Otherwise you end up with duplicate indices, which causes problems further down the line + df = pd.concat([pd.read_csv(path) for path in args.data_paths]).reset_index() + if args.variables: + df = df[args.variables] discover_class = discover_map[args.technique].load() discover = discover_class( diff --git a/causal_testing/discovery/abstract_discovery.py b/causal_testing/discovery/abstract_discovery.py index cb0c049e..64e7404a 100644 --- a/causal_testing/discovery/abstract_discovery.py +++ b/causal_testing/discovery/abstract_discovery.py @@ -166,6 +166,9 @@ def write_dot(self, individual: CausalDAG, output_file: str): else: raise ValueError(f"Invalid test outcome {test['result']}") + individual.remove_node("index") + if "file_index" in individual.nodes: + individual.remove_node("file_index") nx.drawing.nx_pydot.write_dot(individual, output_file) def _json_stub_params(self, outcome: str) -> str: diff --git a/causal_testing/main.py b/causal_testing/main.py index 9d5ed77e..f4c96bf3 100644 --- a/causal_testing/main.py +++ b/causal_testing/main.py @@ -590,6 +590,14 @@ def parse_args(args: Optional[Sequence[str]] = None) -> argparse.Namespace: # Discovery parser_discover = subparsers.add_parser(Command.DISCOVER.value, help="Discover causal structures from data") parser_discover.add_argument("-d", "--data-paths", help="Paths to data files (.csv)", nargs="+", required=True) + parser_discover.add_argument( + "-c", + "--context", + help="Combine data from multiple files into a single DataFrame using a context node rather than concatenation", + action="store_true", + default=False, + required=False, + ) parser_discover.add_argument( "-a", "--alpha", From b72171ce155a64ed60b350216a8f869a87a1eccf Mon Sep 17 00:00:00 2001 From: Jake Newman Date: Tue, 14 Jul 2026 11:50:29 +0100 Subject: [PATCH 2/3] Prevent additional edges --- causal_testing/__main__.py | 13 +++++++------ causal_testing/discovery/abstract_discovery.py | 1 - 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/causal_testing/__main__.py b/causal_testing/__main__.py index d6929e46..ce3c3ae0 100644 --- a/causal_testing/__main__.py +++ b/causal_testing/__main__.py @@ -58,6 +58,7 @@ def main() -> None: kwargs[split[0]] = split[1] logging.info("Discovering causal structure") + exclude_edges = list(nx.nx_pydot.read_dot(args.exclude_edges).edges()) if args.exclude_edges is not None else [] if args.context: dfs = [] @@ -65,23 +66,23 @@ def main() -> None: temp_df = pd.read_csv(path) temp_df['file_index'] = i dfs.append(temp_df) + df = pd.concat(dfs, ignore_index=True) if args.variables: df = df[list(set(args.variables + ['file_index']))] + + exclude_edges.append('".*" -> file_index') else: - # Need to reset index to allow for multiple files having the same index (i.e. starting at zero). - # Otherwise you end up with duplicate indices, which causes problems further down the line - df = pd.concat([pd.read_csv(path) for path in args.data_paths]).reset_index() + df = pd.concat((pd.read_csv(path) for path in args.data_paths), ignore_index=True) + if args.variables: df = df[args.variables] discover_class = discover_map[args.technique].load() discover = discover_class( df=df, - exclude_edges=( - list(nx.nx_pydot.read_dot(args.exclude_edges).edges()) if args.exclude_edges is not None else [] - ), + exclude_edges=exclude_edges, include_edges=( list(nx.nx_pydot.read_dot(args.include_edges).edges()) if args.include_edges is not None else [] ), diff --git a/causal_testing/discovery/abstract_discovery.py b/causal_testing/discovery/abstract_discovery.py index 64e7404a..f9835315 100644 --- a/causal_testing/discovery/abstract_discovery.py +++ b/causal_testing/discovery/abstract_discovery.py @@ -166,7 +166,6 @@ def write_dot(self, individual: CausalDAG, output_file: str): else: raise ValueError(f"Invalid test outcome {test['result']}") - individual.remove_node("index") if "file_index" in individual.nodes: individual.remove_node("file_index") nx.drawing.nx_pydot.write_dot(individual, output_file) From 70a84ec111e0411f44868d100f19b84b7f8dae4f Mon Sep 17 00:00:00 2001 From: Jake Newman <94863788+Jake248Newman@users.noreply.github.com> Date: Tue, 14 Jul 2026 14:53:56 +0100 Subject: [PATCH 3/3] Fix exclude_edges format in __main__.py --- causal_testing/__main__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/causal_testing/__main__.py b/causal_testing/__main__.py index ce3c3ae0..bd49bc6a 100644 --- a/causal_testing/__main__.py +++ b/causal_testing/__main__.py @@ -72,7 +72,7 @@ def main() -> None: if args.variables: df = df[list(set(args.variables + ['file_index']))] - exclude_edges.append('".*" -> file_index') + exclude_edges.append((".*", "file_index")) else: df = pd.concat((pd.read_csv(path) for path in args.data_paths), ignore_index=True)