diff --git a/pyrca/analyzers/base.py b/pyrca/analyzers/base.py index 6d3fe7d..9e9fd8a 100644 --- a/pyrca/analyzers/base.py +++ b/pyrca/analyzers/base.py @@ -1,11 +1,12 @@ # -# Copyright (c) 2023 salesforce.com, inc. +# Copyright (c) 2026 salesforce.com, inc. # All rights reserved. # SPDX-License-Identifier: BSD-3-Clause # For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause# """Base classes for all RCA algorithms""" from abc import abstractmethod from dataclasses import dataclass, field, asdict +import pandas as pd from pyrca.base import BaseModel from pyrca.utils.logger import get_logger @@ -52,6 +53,15 @@ class BaseRCA(BaseModel): def __init__(self): self.logger = get_logger(self.__class__.__name__) + @staticmethod + def _load_csv_graph(path): + graph = pd.read_csv(path) + if graph.shape[1] == graph.shape[0] + 1: + graph = pd.read_csv(path, index_col=0) + elif graph.shape[1] == graph.shape[0]: + graph.index = graph.columns + return graph + @abstractmethod def train(self, **kwargs): """ diff --git a/pyrca/analyzers/bayesian.py b/pyrca/analyzers/bayesian.py index 55023ff..75c28ff 100644 --- a/pyrca/analyzers/bayesian.py +++ b/pyrca/analyzers/bayesian.py @@ -1,5 +1,5 @@ # -# Copyright (c) 2023 salesforce.com, inc. +# Copyright (c) 2026 salesforce.com, inc. # All rights reserved. # SPDX-License-Identifier: BSD-3-Clause # For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause# @@ -32,6 +32,7 @@ class BayesianNetworkConfig(BaseConfig): :param graph: The adjacency matrix of the causal graph, which can be a pandas dataframe or a file path of a CSV file or a pickled file. + CSV row labels may be included in the first column; otherwise rows follow the column order. :param sigmas: Specific sigmas other than ``default_sigma`` for certain variables. This parameter is for constructing training data (only used when ``detector`` in the ``train`` function is not set). :param default_sigma: The default sigma value for computing the detection. This parameter is @@ -67,7 +68,7 @@ def __init__(self, config: BayesianNetworkConfig): self.config = config if isinstance(config.graph, str): if config.graph.endswith(".csv"): - self.graph = pd.read_csv(config.graph) + self.graph = self._load_csv_graph(config.graph) elif config.graph.endswith(".pkl"): with open(config.graph, "rb") as f: self.graph = pickle.load(f) diff --git a/pyrca/analyzers/ht.py b/pyrca/analyzers/ht.py index e60e982..95b7df8 100644 --- a/pyrca/analyzers/ht.py +++ b/pyrca/analyzers/ht.py @@ -1,5 +1,5 @@ # -# Copyright (c) 2023 salesforce.com, inc. +# Copyright (c) 2026 salesforce.com, inc. # All rights reserved. # SPDX-License-Identifier: BSD-3-Clause # For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause# @@ -26,6 +26,7 @@ class HTConfig(BaseConfig): :param graph: The adjacency matrix of the causal graphs, which can be a pandas dataframe or a file path of a CSV file or a pickled file. + CSV row labels may be included in the first column; otherwise rows follow the column order. :param aggregator: The function for aggregating the node score from all the abnormal data. :param root_cause_top_k: The maximum number of root causes in the results. """ @@ -49,7 +50,7 @@ def __init__(self, config: HTConfig): self.config = config if isinstance(config.graph, str): if config.graph.endswith(".csv"): - graph = pd.read_csv(config.graph) + graph = self._load_csv_graph(config.graph) elif config.graph.endswith(".pkl"): with open(config.graph, "rb") as f: graph = pickle.load(f) diff --git a/pyrca/analyzers/random_walk.py b/pyrca/analyzers/random_walk.py index fab9b7a..8579fd3 100644 --- a/pyrca/analyzers/random_walk.py +++ b/pyrca/analyzers/random_walk.py @@ -1,5 +1,5 @@ # -# Copyright (c) 2023 salesforce.com, inc. +# Copyright (c) 2026 salesforce.com, inc. # All rights reserved. # SPDX-License-Identifier: BSD-3-Clause # For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause# @@ -25,6 +25,7 @@ class RandomWalkConfig(BaseConfig): :param graph: The adjacency matrix of the causal graph, which can be a pandas dataframe or a file path of a CSV file or a pickled file. + CSV row labels may be included in the first column; otherwise rows follow the column order. :param use_partial_corr: Whether to use partial correlation when computing edge weights. :param rho: The weight from a "cause" node to a "result" node. :param num_steps: The number of random walk steps in each run. @@ -52,7 +53,7 @@ def __init__(self, config: RandomWalkConfig): self.config = config if isinstance(config.graph, str): if config.graph.endswith(".csv"): - graph = pd.read_csv(config.graph) + graph = self._load_csv_graph(config.graph) elif config.graph.endswith(".pkl"): with open(config.graph, "rb") as f: graph = pickle.load(f) diff --git a/tests/analyzers/test_csv_graph.py b/tests/analyzers/test_csv_graph.py new file mode 100644 index 0000000..afbf045 --- /dev/null +++ b/tests/analyzers/test_csv_graph.py @@ -0,0 +1,30 @@ +# +# Copyright (c) 2026 salesforce.com, inc. +# All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause +# For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause# +import pandas as pd +import pytest + +from pyrca.analyzers.bayesian import BayesianNetwork +from pyrca.analyzers.ht import HT +from pyrca.analyzers.random_walk import RandomWalk + + +@pytest.mark.parametrize("analyzer", [BayesianNetwork, HT, RandomWalk]) +@pytest.mark.parametrize("index_name,write_index", [(None, True), ("metric", True), (None, False)]) +def test_analyzer_loads_csv_adjacency_matrix(tmp_path, analyzer, index_name, write_index): + graph = pd.DataFrame([[0, 1], [0, 0]], index=["cause", "effect"], columns=["cause", "effect"]) + graph.index.name = index_name + path = tmp_path / "graph.csv" + graph.to_csv(path, index=write_index) + + model = analyzer(analyzer.config_class(graph=str(path))) + if analyzer is BayesianNetwork: + pd.testing.assert_frame_equal(model.graph, graph) + actual_graph = model.bayesian_model + else: + pd.testing.assert_frame_equal(model.adjacency_mat, graph) + actual_graph = model.graph + assert set(actual_graph.nodes) == {"cause", "effect"} + assert set(actual_graph.edges) == {("cause", "effect")}