Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion pyrca/analyzers/base.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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):
"""
Expand Down
5 changes: 3 additions & 2 deletions pyrca/analyzers/bayesian.py
Original file line number Diff line number Diff line change
@@ -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#
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
5 changes: 3 additions & 2 deletions pyrca/analyzers/ht.py
Original file line number Diff line number Diff line change
@@ -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#
Expand All @@ -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.
"""
Expand All @@ -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)
Expand Down
5 changes: 3 additions & 2 deletions pyrca/analyzers/random_walk.py
Original file line number Diff line number Diff line change
@@ -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#
Expand All @@ -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.
Expand Down Expand Up @@ -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)
Expand Down
30 changes: 30 additions & 0 deletions tests/analyzers/test_csv_graph.py
Original file line number Diff line number Diff line change
@@ -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")}