File size: 3,344 Bytes
8e874f5 | 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 | """Retrievers module - modular graph retrieval algorithms.
This module provides pluggable retrieval algorithms for knowledge graph traversal.
Similar to how neural network frameworks organize optimizers (Adam, SGD, etc.),
this module organizes graph traversal algorithms.
Available Retrievers:
- FlowDiffusionRetriever: Query-Aware Weighted Flow Diffusion algorithm
Usage:
# Direct import
from src.retrievers import FlowDiffusionRetriever
retriever = FlowDiffusionRetriever(
graph=my_graph,
source_node="node_a",
target_node="node_b",
confidence=0.5
)
result = retriever.retrieve()
# Or use the registry
from src.retrievers import get_retriever
retriever = get_retriever(
"flow_diffusion",
graph=my_graph,
source_node="node_a",
target_node="node_b"
)
Adding Custom Retrievers:
from src.retrievers import BaseRetriever, register_retriever
class MyCustomRetriever(BaseRetriever):
def retrieve(self, source_node, target_node=None, **kwargs):
# Implementation
pass
def find_path(self, source, target):
# Implementation
pass
register_retriever("my_custom", MyCustomRetriever)
"""
from .base import BaseRetriever, RetrieverResult
from .flow_diffusion import FlowDiffusionRetriever, QueryAwareWeightedFlowDiffusion
# Registry for retriever classes
RETRIEVER_REGISTRY = {
"flow_diffusion": FlowDiffusionRetriever,
"qafd": FlowDiffusionRetriever, # Alias
}
def get_retriever(name: str, **kwargs) -> BaseRetriever:
"""Factory function to get a retriever by name.
Args:
name: Retriever name (e.g., "flow_diffusion", "qafd")
**kwargs: Arguments to pass to the retriever constructor
Returns:
Instantiated retriever
Raises:
ValueError: If retriever name is not found
Example:
retriever = get_retriever(
"flow_diffusion",
graph=my_graph,
source_node="node_a",
target_node="node_b",
confidence=0.5
)
"""
if name not in RETRIEVER_REGISTRY:
available = ", ".join(RETRIEVER_REGISTRY.keys())
raise ValueError(f"Unknown retriever: '{name}'. Available retrievers: {available}")
return RETRIEVER_REGISTRY[name](**kwargs)
def register_retriever(name: str, cls: type):
"""Register a custom retriever class.
Args:
name: Name to register the retriever under
cls: Retriever class (should inherit from BaseRetriever)
Example:
register_retriever("my_retriever", MyRetrieverClass)
"""
if not issubclass(cls, BaseRetriever):
raise TypeError(f"Retriever class must inherit from BaseRetriever, got {cls}")
RETRIEVER_REGISTRY[name] = cls
def list_retrievers() -> list:
"""List all available retriever names.
Returns:
List of registered retriever names
"""
return list(RETRIEVER_REGISTRY.keys())
__all__ = [
# Base classes
"BaseRetriever",
"RetrieverResult",
# Retriever implementations
"FlowDiffusionRetriever",
"QueryAwareWeightedFlowDiffusion", # Backward compatibility alias
# Registry functions
"get_retriever",
"register_retriever",
"list_retrievers",
"RETRIEVER_REGISTRY",
]
|