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 )