Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/evaluation/retrieval.py: 82%

40 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 07:19 +0800

1"""Retrieval-based evaluation metrics.""" 

2 

3from __future__ import annotations 

4 

5from typing import Any 

6 

7from lexigram.ai.rag.evaluation.base import EvaluatorBase 

8from lexigram.ai.rag.evaluation.types import EvaluationResult, MetricType 

9 

10 

11class RetrievalPrecisionEvaluator(EvaluatorBase): 

12 """Evaluates retrieval precision. 

13 

14 Measures the fraction of retrieved documents that are relevant. 

15 Requires ground truth relevant document IDs. 

16 """ 

17 

18 def __init__(self) -> None: 

19 """Initialize retrieval precision evaluator.""" 

20 super().__init__("retrieval_precision") 

21 

22 async def evaluate( 

23 self, 

24 query: str, 

25 retrieved_docs: list[Any], 

26 generated_answer: str, 

27 reference_answer: str | None = None, 

28 **kwargs, 

29 ) -> EvaluationResult: 

30 """Evaluate retrieval precision. 

31 

32 Args: 

33 query: The query. 

34 retrieved_docs: Retrieved documents (should have 'id' or be IDs). 

35 generated_answer: Generated answer (not used). 

36 reference_answer: Not used. 

37 **kwargs: Must contain 'relevant_doc_ids' (set of relevant IDs). 

38 

39 Returns: 

40 Precision score. 

41 """ 

42 relevant_doc_ids = kwargs.get("relevant_doc_ids", set()) 

43 

44 if not retrieved_docs: 

45 return EvaluationResult( 

46 metric_type=MetricType.RETRIEVAL_PRECISION, 

47 score=0.0, 

48 details={"reason": "No documents retrieved"}, 

49 ) 

50 

51 # Extract IDs from retrieved docs 

52 retrieved_ids = set() 

53 for doc in retrieved_docs: 

54 if isinstance(doc, dict): 

55 retrieved_ids.add(doc.get("id", doc.get("doc_id"))) 

56 elif hasattr(doc, "id"): 

57 retrieved_ids.add(doc.id) 

58 else: 

59 retrieved_ids.add(str(doc)) 

60 

61 # Calculate precision 

62 if not retrieved_ids: 

63 precision = 0.0 

64 else: 

65 relevant_retrieved = retrieved_ids & relevant_doc_ids 

66 precision = len(relevant_retrieved) / len(retrieved_ids) 

67 

68 return EvaluationResult( 

69 metric_type=MetricType.RETRIEVAL_PRECISION, 

70 score=precision, 

71 details={ 

72 "retrieved_count": len(retrieved_ids), 

73 "relevant_count": len(relevant_doc_ids), 

74 "relevant_retrieved": len(retrieved_ids & relevant_doc_ids), 

75 }, 

76 ) 

77 

78 

79class RetrievalRecallEvaluator(EvaluatorBase): 

80 """Evaluates retrieval recall. 

81 

82 Measures the fraction of relevant documents that were retrieved. 

83 Requires ground truth relevant document IDs. 

84 """ 

85 

86 def __init__(self) -> None: 

87 """Initialize retrieval recall evaluator.""" 

88 super().__init__("retrieval_recall") 

89 

90 async def evaluate( 

91 self, 

92 query: str, 

93 retrieved_docs: list[Any], 

94 generated_answer: str, 

95 reference_answer: str | None = None, 

96 **kwargs, 

97 ) -> EvaluationResult: 

98 """Evaluate retrieval recall. 

99 

100 Args: 

101 query: The query. 

102 retrieved_docs: Retrieved documents. 

103 generated_answer: Generated answer (not used). 

104 reference_answer: Not used. 

105 **kwargs: Must contain 'relevant_doc_ids'. 

106 

107 Returns: 

108 Recall score. 

109 """ 

110 relevant_doc_ids = kwargs.get("relevant_doc_ids", set()) 

111 

112 if not relevant_doc_ids: 

113 return EvaluationResult( 

114 metric_type=MetricType.RETRIEVAL_RECALL, 

115 score=0.0, 

116 details={"reason": "No relevant documents specified"}, 

117 ) 

118 

119 # Extract IDs 

120 retrieved_ids = set() 

121 for doc in retrieved_docs: 

122 if isinstance(doc, dict): 

123 retrieved_ids.add(doc.get("id", doc.get("doc_id"))) 

124 elif hasattr(doc, "id"): 

125 retrieved_ids.add(doc.id) 

126 else: 

127 retrieved_ids.add(str(doc)) 

128 

129 # Calculate recall 

130 relevant_retrieved = retrieved_ids & relevant_doc_ids 

131 recall = len(relevant_retrieved) / len(relevant_doc_ids) 

132 

133 return EvaluationResult( 

134 metric_type=MetricType.RETRIEVAL_RECALL, 

135 score=recall, 

136 details={ 

137 "retrieved_count": len(retrieved_ids), 

138 "relevant_count": len(relevant_doc_ids), 

139 "relevant_retrieved": len(relevant_retrieved), 

140 }, 

141 )