Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-rag/src/lexigram/ai/rag/hyde/generators/weighted.py: 96%

26 statements  

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

1"""Weighted HyDE generator implementation.""" 

2 

3from __future__ import annotations 

4 

5from typing import Any 

6 

7from lexigram.ai.rag.hyde.base import AbstractHyDEGenerator 

8from lexigram.ai.rag.hyde.protocols import EmbeddingClientProtocol 

9from lexigram.ai.rag.hyde.types import HyDEResult, HyDEStrategy, HypotheticalDocument 

10from lexigram.contracts import ( 

11 LLMClientProtocol, 

12) 

13 

14 

15class WeightedHyDEGenerator(AbstractHyDEGenerator): 

16 """Generator with weighted aggregation of multiple documents.""" 

17 

18 def __init__( 

19 self, 

20 llm_client: LLMClientProtocol, 

21 embedding_client: EmbeddingClientProtocol, 

22 temperature: float = 0.8, 

23 max_tokens: int = 150, 

24 default_num_documents: int = 3, 

25 confidence_decay: float = 0.7, 

26 ): 

27 """Initialize weighted HyDE generator. 

28 

29 Args: 

30 llm_client: Client for generating hypothetical documents 

31 embedding_client: Client for generating embeddings (required) 

32 temperature: Sampling temperature 

33 max_tokens: Maximum tokens per document 

34 default_num_documents: Default number of documents 

35 confidence_decay: Decay factor for subsequent documents 

36 """ 

37 super().__init__(llm_client, embedding_client) 

38 self.temperature = temperature 

39 self.max_tokens = max_tokens 

40 self.default_num_documents = default_num_documents 

41 self.confidence_decay = confidence_decay 

42 

43 async def generate( 

44 self, 

45 query: str, 

46 num_documents: int | None = None, 

47 **kwargs: Any, 

48 ) -> HyDEResult: 

49 """Generate weighted hypothetical documents. 

50 

51 Args: 

52 query: User query 

53 num_documents: Number of documents 

54 **kwargs: Additional parameters 

55 

56 Returns: 

57 HyDE result with weighted aggregation 

58 """ 

59 if num_documents is None: 

60 num_documents = self.default_num_documents 

61 

62 # Generate multiple documents 

63 documents = [] 

64 for i in range(num_documents): 

65 content = await self._generate_single_document( 

66 query, 

67 temperature=self.temperature, 

68 max_tokens=self.max_tokens, 

69 **kwargs, 

70 ) 

71 

72 # Exponential confidence decay 

73 confidence = self.confidence_decay**i 

74 

75 doc = HypotheticalDocument( 

76 content=content, 

77 query=query, 

78 confidence=confidence, 

79 metadata={ 

80 "index": i, 

81 "temperature": self.temperature, 

82 "max_tokens": self.max_tokens, 

83 }, 

84 ) 

85 documents.append(doc) 

86 

87 # Generate embeddings (required for weighted strategy) 

88 embeddings = await self._embed_documents(documents) 

89 

90 # Weighted aggregation based on confidence 

91 weights = [doc.confidence for doc in documents] 

92 aggregated_embedding = self._aggregate_embeddings(embeddings, weights) 

93 

94 return HyDEResult( 

95 query=query, 

96 hypothetical_docs=documents, 

97 strategy=HyDEStrategy.WEIGHTED, 

98 aggregated_embedding=aggregated_embedding, 

99 metadata={ 

100 "temperature": self.temperature, 

101 "max_tokens": self.max_tokens, 

102 "num_documents": num_documents, 

103 "confidence_decay": self.confidence_decay, 

104 "weights": weights, 

105 }, 

106 )