Spaces:
Running
Running
File size: 5,664 Bytes
42fb3af eb40daa 42fb3af eb40daa 42fb3af eb40daa 42fb3af eb40daa 42fb3af eb40daa 42fb3af eb40daa 42fb3af | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 | """Graph traversal: entity expansion, trial lookup, evidence retrieval."""
from __future__ import annotations
import networkx as nx
from config import KG_EXPANSION_HOPS, KG_MIN_EDGE_CONFIDENCE
from extraction.normalizer import normalize_entity, _GENE_ALIASES, _COMPOUND_ALIASES
from logging_config import get_logger
_logger = get_logger("graph.query")
def expand_query_entities(
G: nx.DiGraph,
entity_names: list[str],
max_hops: int = KG_EXPANSION_HOPS,
) -> list[str]:
"""
Map entity names to graph node IDs, then BFS-expand up to max_hops.
Returns a deduplicated list of entity display_names for RAG search.
Only traverses edges above KG_MIN_EDGE_CONFIDENCE.
"""
if not G or not entity_names:
return entity_names
# Phase 1: map names to canonical node IDs
seed_nodes: set[str] = set()
for name in entity_names:
matched = _find_node(G, name)
if matched:
seed_nodes.update(matched)
else:
_logger.debug(f"No graph node for query entity: {name!r}")
if not seed_nodes:
return entity_names # fall through to lexical RAG search
# Phase 2: BFS expansion
frontier = set(seed_nodes)
expanded = set(seed_nodes)
for _ in range(max_hops):
next_frontier: set[str] = set()
for node in frontier:
for neighbor in list(G.successors(node)) + list(G.predecessors(node)):
if neighbor in expanded:
continue
# Only follow confident edges
edge_data = G.get_edge_data(node, neighbor) or G.get_edge_data(neighbor, node) or {}
if edge_data.get("confidence", 1.0) >= KG_MIN_EDGE_CONFIDENCE:
# Skip trial nodes — they inflate retrieval noise
if G.nodes[neighbor].get("type") != "ClinicalTrial":
next_frontier.add(neighbor)
expanded.update(next_frontier)
frontier = next_frontier
# Phase 3: convert canonical IDs back to display names for RAG text queries
display_names: list[str] = []
seen: set[str] = set()
for node_id in expanded:
if node_id.startswith("trial:"):
continue
name = G.nodes[node_id].get("display_name", node_id.split(":", 1)[-1])
if name not in seen:
display_names.append(name)
seen.add(name)
_logger.info(
f"KG expansion: {len(entity_names)} query entities → {len(display_names)} expanded",
extra={"data": {"seeds": list(seed_nodes), "expanded_count": len(expanded)}},
)
return display_names
_STATUS_RANK = {"RECRUITING": 0, "ACTIVE_NOT_RECRUITING": 1, "NOT_YET_RECRUITING": 2, "COMPLETED": 3}
def find_trials_for_entities(
G: nx.DiGraph,
entity_names: list[str],
max_trials: int = 10,
) -> list[dict]:
"""Return clinical trials linked to the given entity names.
Collects all matches, scores by number of linked entities, sorts by
status (RECRUITING first) then score, and returns the top max_trials.
"""
if not G or not entity_names:
return []
target_nodes: set[str] = set()
for name in entity_names:
matched = _find_node(G, name)
target_nodes.update(matched)
# score[nct_id] = number of query entities this trial links to
scores: dict[str, int] = {}
meta: dict[str, dict] = {}
for node_id in target_nodes:
for pred in G.predecessors(node_id):
if not pred.startswith("trial:"):
continue
nct_id = G.nodes[pred].get("nct_id", "")
if not nct_id:
continue
scores[nct_id] = scores.get(nct_id, 0) + 1
if nct_id not in meta:
status = G.nodes[pred].get("status", "")
meta[nct_id] = {
"nct_id": nct_id,
"title": G.nodes[pred].get("display_name", ""),
"phase": G.nodes[pred].get("phase", ""),
"status": status,
"url": G.nodes[pred].get("url", ""),
"_status_rank": _STATUS_RANK.get(status, 9),
}
ranked = sorted(
meta.values(),
key=lambda t: (t["_status_rank"], -scores[t["nct_id"]]),
)
for t in ranked:
del t["_status_rank"]
return ranked[:max_trials]
def get_entity_evidence(G: nx.DiGraph, canonical_id: str) -> dict:
"""Return node attributes + connected entity names for a canonical ID."""
if not G.has_node(canonical_id):
return {}
attrs = dict(G.nodes[canonical_id])
attrs["neighbors"] = [
{
"id": n,
"display_name": G.nodes[n].get("display_name", n),
"relation": G.get_edge_data(canonical_id, n, {}).get("relation_type", ""),
}
for n in G.successors(canonical_id)
if G.nodes[n].get("type") != "ClinicalTrial"
]
return attrs
def _find_node(G: nx.DiGraph, name: str) -> list[str]:
"""Map a raw entity name to zero or more graph node IDs."""
hits: list[str] = []
# 1. Try each entity type prefix
for etype in ("Gene", "Protein", "Compound", "Mechanism", "Pathway", "Phenotype"):
candidate = normalize_entity(name, etype)
if G.has_node(candidate):
hits.append(candidate)
if hits:
return hits
# 2. Case-insensitive display_name match
name_lower = name.lower()
for node_id, data in G.nodes(data=True):
display = data.get("display_name", "").lower()
if display == name_lower or name_lower in display:
hits.append(node_id)
return hits
|