mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
fix(graphrag): rank without scipy so graph retrieval stops falling back
networkx.pagerank delegates to a scipy implementation, and scipy is not a DocsGPT dependency — it only arrives transitively through the optional docling extra. In a default install every graph retrieval raised ModuleNotFoundError inside _ppr_scores, hit the per-source except, and degraded to ClassicRAG: the graph was built and paid for, then never used, with one ERROR line per source per query as the only signal. Rank with a local power iteration over the same row-normalized transition matrix: undirected edges normalized per endpoint, dangling nodes redistributed along the restart vector, and the restart vector normalized across the nodes the subgraph actually holds so seed mass cannot leak. Parity with networkx is asserted while scipy happens to be installed in the test env, and the retrieval path is exercised with the import blocked.
This commit is contained in:
1 parent
24bcbd938e
commit
94a33c4781
2 files changed
+226
-5
No files matched your search
@@ -1,9 +1,13 @@
|
||||
"""GraphRAG local retriever — Personalized PageRank over a per-source graph.
|
||||
|
||||
Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> networkx
|
||||
Personalized PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR
|
||||
mass -> shared token budget. No LLM call at query time beyond the (optional,
|
||||
reused) rephrase.
|
||||
Rephrased query -> entity-name NN seeds -> bounded 1-2-hop fetch -> Personalized
|
||||
PageRank (IDF-down-weighted hubs) -> chunks ranked by landed PPR mass -> shared
|
||||
token budget. No LLM call at query time beyond the (optional, reused) rephrase.
|
||||
|
||||
``networkx`` supplies the graph structure, but the ranking is the local power
|
||||
iteration in :func:`_personalized_pagerank`: ``nx.pagerank`` delegates to scipy,
|
||||
which DocsGPT does not depend on, so calling it turned every graph retrieval
|
||||
into a silent ClassicRAG fallback.
|
||||
|
||||
Composes :class:`ClassicRAG` rather than subclassing: PPR doesn't fit the
|
||||
``_fetch_candidates`` hook, but the composed instance supplies the rephrase, the
|
||||
@@ -40,6 +44,99 @@ def _idf(doc_freq: Any) -> float:
|
||||
return 1.0 / math.log(1.0 + max(int(doc_freq or 0), 0) + 1.0)
|
||||
|
||||
|
||||
def _restart_vector(nodes: List[Any], personalization: Dict[Any, float] | None) -> Dict[Any, float]:
|
||||
"""Normalized restart distribution over ``nodes``.
|
||||
|
||||
Weights are clamped at zero (a cosine distance above 1 yields a negative
|
||||
seed weight) and normalized across the nodes actually present in the graph,
|
||||
so restart mass can never leak to a node the subgraph does not contain. An
|
||||
absent or all-zero personalization collapses to a uniform restart.
|
||||
"""
|
||||
if personalization:
|
||||
weights = {
|
||||
node: max(float(personalization.get(node, 0.0) or 0.0), 0.0)
|
||||
for node in nodes
|
||||
}
|
||||
total = sum(weights.values())
|
||||
if total > 0:
|
||||
return {node: weight / total for node, weight in weights.items()}
|
||||
uniform = 1.0 / len(nodes)
|
||||
return {node: uniform for node in nodes}
|
||||
|
||||
|
||||
def _personalized_pagerank(
|
||||
graph: nx.Graph,
|
||||
personalization: Dict[Any, float] | None = None,
|
||||
*,
|
||||
weight: str = "weight",
|
||||
alpha: float = 0.85,
|
||||
max_iter: int = 100,
|
||||
tol: float = 1.0e-6,
|
||||
) -> Dict[Any, float]:
|
||||
"""Personalized PageRank by power iteration — no scipy.
|
||||
|
||||
``networkx.pagerank`` delegates to a scipy implementation, and scipy is not
|
||||
a DocsGPT dependency: in a default install the import raises and every graph
|
||||
retrieval silently degrades to the ClassicRAG fallback. This is the same
|
||||
algorithm over the same row-normalized transition matrix, so the ranking is
|
||||
unchanged where scipy happens to be installed.
|
||||
|
||||
Args:
|
||||
graph: Undirected graph whose edges may carry a ``weight`` attribute.
|
||||
personalization: Node -> restart weight; ``None`` means uniform.
|
||||
weight: Edge attribute holding the weight.
|
||||
alpha: Damping factor.
|
||||
max_iter: Iteration cap. The last iterate is returned if it is hit —
|
||||
retrieval degrades to a slightly less converged ranking rather than
|
||||
raising, which is what the library does.
|
||||
tol: Convergence tolerance; iteration stops below ``len(graph) * tol``.
|
||||
|
||||
Returns:
|
||||
Node -> PageRank mass, summing to ~1.0. Empty dict for an empty graph.
|
||||
"""
|
||||
nodes = list(graph.nodes)
|
||||
node_count = len(nodes)
|
||||
if node_count == 0:
|
||||
return {}
|
||||
|
||||
restart = _restart_vector(nodes, personalization)
|
||||
|
||||
# Row-normalized transitions. An undirected edge is traversable from both
|
||||
# endpoints, so each node normalizes over its own incident weights.
|
||||
transitions: Dict[Any, List[Any]] = {}
|
||||
for node in nodes:
|
||||
neighbors = []
|
||||
total = 0.0
|
||||
for neighbor, data in graph[node].items():
|
||||
edge_weight = float(data.get(weight, 1.0) or 1.0)
|
||||
if edge_weight <= 0:
|
||||
continue
|
||||
neighbors.append((neighbor, edge_weight))
|
||||
total += edge_weight
|
||||
transitions[node] = (
|
||||
[(n, w / total) for n, w in neighbors] if total > 0 else []
|
||||
)
|
||||
|
||||
# A node with no usable edge is dangling: its mass would vanish each pass,
|
||||
# so it is redistributed along the restart vector instead.
|
||||
dangling = [node for node in nodes if not transitions[node]]
|
||||
|
||||
ranks = {node: 1.0 / node_count for node in nodes}
|
||||
for _ in range(max_iter):
|
||||
previous = ranks
|
||||
ranks = dict.fromkeys(nodes, 0.0)
|
||||
leaked = alpha * sum(previous[node] for node in dangling)
|
||||
for node in nodes:
|
||||
share = alpha * previous[node]
|
||||
for neighbor, transition in transitions[node]:
|
||||
ranks[neighbor] += share * transition
|
||||
for node in nodes:
|
||||
ranks[node] += (leaked + 1.0 - alpha) * restart[node]
|
||||
if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol:
|
||||
break
|
||||
return ranks
|
||||
|
||||
|
||||
class GraphRAGRetriever(BaseRetriever):
|
||||
"""Per-source PPR retriever; falls back to ClassicRAG when a source has no graph."""
|
||||
|
||||
@@ -117,7 +214,9 @@ class GraphRAGRetriever(BaseRetriever):
|
||||
if not any(personalization.values()):
|
||||
personalization = None
|
||||
|
||||
ranks = nx.pagerank(graph, personalization=personalization, weight="weight")
|
||||
ranks = _personalized_pagerank(
|
||||
graph, personalization=personalization, weight="weight"
|
||||
)
|
||||
return {
|
||||
node: rank * _idf(graph.nodes[node].get("doc_freq", 0))
|
||||
for node, rank in ranks.items()
|
||||
|
||||
@@ -800,3 +800,125 @@ class TestGraphRAGBatching:
|
||||
assert rag._get_data() == []
|
||||
mock_store_cls.assert_not_called()
|
||||
rag._classic._get_data.assert_not_called()
|
||||
|
||||
|
||||
# ── Personalized PageRank without scipy ──────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _no_scipy(monkeypatch):
|
||||
"""Make ``import scipy`` fail, as it does in a default install.
|
||||
|
||||
``scipy`` is not a DocsGPT dependency — it only reaches this test env
|
||||
through the optional docling extra. ``networkx.pagerank`` delegates to its
|
||||
scipy implementation, so ranking must not go through it.
|
||||
"""
|
||||
import sys
|
||||
|
||||
for name in [m for m in list(sys.modules) if m == "scipy" or m.startswith("scipy.")]:
|
||||
monkeypatch.delitem(sys.modules, name)
|
||||
monkeypatch.setitem(sys.modules, "scipy", None)
|
||||
|
||||
|
||||
def _chain_graph():
|
||||
"""Weighted chain a-b-c-d plus a heavier shortcut a-d."""
|
||||
import networkx as nx
|
||||
|
||||
graph = nx.Graph()
|
||||
graph.add_weighted_edges_from(
|
||||
[("a", "b", 1.0), ("b", "c", 2.0), ("c", "d", 1.0), ("a", "d", 0.5)]
|
||||
)
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestPersonalizedPageRankWithoutScipy:
|
||||
def test_ranking_runs_when_scipy_is_missing(self, _no_scipy):
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
ranks = _personalized_pagerank(
|
||||
graph, personalization={"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
|
||||
)
|
||||
|
||||
assert set(ranks) == {"a", "b", "c", "d"}
|
||||
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
|
||||
assert all(rank > 0 for rank in ranks.values())
|
||||
# Pinned from the parity test below, which runs the library
|
||||
# implementation over the same graph while scipy is installed here.
|
||||
assert sorted(ranks, key=ranks.get, reverse=True) == ["b", "a", "c", "d"]
|
||||
# The seed outranks the node furthest from it along the heavy path.
|
||||
assert ranks["a"] > ranks["d"]
|
||||
|
||||
def test_matches_networkx_within_tolerance(self):
|
||||
"""Parity with the library implementation, while it is installed here."""
|
||||
import networkx as nx
|
||||
|
||||
pytest.importorskip("scipy")
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
personalization = {"a": 1.0, "b": 0.0, "c": 0.0, "d": 0.0}
|
||||
|
||||
ours = _personalized_pagerank(graph, personalization=personalization)
|
||||
theirs = nx.pagerank(graph, personalization=personalization, weight="weight")
|
||||
|
||||
for node in theirs:
|
||||
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
|
||||
|
||||
def test_uniform_personalization_when_none(self):
|
||||
pytest.importorskip("scipy")
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = _chain_graph()
|
||||
ours = _personalized_pagerank(graph, personalization=None)
|
||||
theirs = nx.pagerank(graph, personalization=None, weight="weight")
|
||||
|
||||
for node in theirs:
|
||||
assert ours[node] == pytest.approx(theirs[node], abs=1e-6)
|
||||
|
||||
def test_isolated_node_still_gets_mass(self):
|
||||
"""A node with no edges is dangling; its mass must not vanish."""
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
graph = nx.Graph()
|
||||
graph.add_edge("a", "b", weight=1.0)
|
||||
graph.add_node("lonely")
|
||||
|
||||
ranks = _personalized_pagerank(graph, personalization=None)
|
||||
|
||||
assert ranks["lonely"] > 0
|
||||
assert sum(ranks.values()) == pytest.approx(1.0, abs=1e-6)
|
||||
|
||||
def test_empty_graph_returns_empty(self):
|
||||
import networkx as nx
|
||||
|
||||
from docsgpt.retriever.graph_rag import _personalized_pagerank
|
||||
|
||||
assert _personalized_pagerank(nx.Graph(), personalization=None) == {}
|
||||
|
||||
@patch("docsgpt.retriever.graph_rag.num_tokens_from_string", return_value=10)
|
||||
@patch("docsgpt.retriever.graph_rag.GraphStore")
|
||||
@patch("docsgpt.retriever.graph_rag.graphrag_available", return_value=True)
|
||||
def test_graph_retrieval_does_not_fall_back_without_scipy(
|
||||
self, _avail, mock_store_cls, _tok, _patch_llm_creator, _patch_embed, _no_scipy
|
||||
):
|
||||
"""The whole PPR path runs with scipy absent — no ClassicRAG fallback."""
|
||||
nodes = [{"id": "n1", "doc_freq": 1}, {"id": "n2", "doc_freq": 1}]
|
||||
edges = [{"src_node_id": "n1", "dst_node_id": "n2", "weight": 1.0}]
|
||||
node_chunks = {"n1": ["c1"], "n2": ["c2"]}
|
||||
chunk_texts = {"c1": "near", "c2": "far"}
|
||||
seed_rows = [{"id": "n1", "distance": 0.0}]
|
||||
store = _store_with_graph(nodes, edges, node_chunks, chunk_texts, seed_rows)
|
||||
mock_store_cls.return_value = store
|
||||
|
||||
rag = _make_retriever(chunks=2)
|
||||
rag._classic_for_sources = Mock(side_effect=AssertionError("fell back"))
|
||||
|
||||
docs = rag._get_data()
|
||||
|
||||
assert [doc["text"] for doc in docs] == ["near", "far"]
|
||||
Reference in new issue
Block a user