151 lines
5.2 KiB
Python
151 lines
5.2 KiB
Python
"""
|
|
Query transformer for RAG system
|
|
Implements three query transformation techniques:
|
|
1. Query rewriting: Make queries more specific and detailed
|
|
2. Backward prompt generation: Generate broader queries for context retrieval
|
|
3. Sub-query decomposition: Break down complex queries into simpler components
|
|
"""
|
|
from typing import List, Dict, Optional, Tuple
|
|
from loguru import logger
|
|
from config import settings
|
|
from llama_index.llms.ollama import Ollama
|
|
from utils.prompt.query_trans import REWRITE_PROMPT, BACKWARD_PROMPT, DECOMPOSE_PROMPT
|
|
|
|
|
|
class QueryTransformer:
|
|
"""
|
|
Query transformer for RAG system
|
|
"""
|
|
|
|
def __init__(self, llm: Ollama):
|
|
"""
|
|
Initialize query transformer
|
|
|
|
Args:
|
|
llm: Ollama LLM instance for generating transformations
|
|
"""
|
|
self.llm = llm
|
|
|
|
async def rewrite_query(self, query: str, history: str = "") -> str:
|
|
"""
|
|
Rewrite query to be more specific and detailed
|
|
|
|
Args:
|
|
query: Original user query
|
|
history: Chat history (optional)
|
|
|
|
Returns:
|
|
Rewritten query
|
|
"""
|
|
try:
|
|
prompt = REWRITE_PROMPT.format(query=query, history=history)
|
|
|
|
response = await self.llm.acomplete(prompt=prompt)
|
|
rewritten_query = response.text.strip()
|
|
logger.info(f"Rewritten query: {rewritten_query}")
|
|
return rewritten_query
|
|
except Exception as e:
|
|
logger.error(f"Error in query rewriting: {e}")
|
|
return query
|
|
|
|
async def backward_query(self, query: str, history: str = "") -> str:
|
|
"""
|
|
Generate broader query for context retrieval
|
|
|
|
Args:
|
|
query: Original user query
|
|
history: Chat history (optional)
|
|
|
|
Returns:
|
|
Broader query for context retrieval
|
|
"""
|
|
try:
|
|
prompt = BACKWARD_PROMPT.format(query=query, history=history)
|
|
|
|
response = await self.llm.acomplete(prompt=prompt)
|
|
backward_query = response.text.strip()
|
|
logger.info(f"Backward query: {backward_query}")
|
|
return backward_query
|
|
except Exception as e:
|
|
logger.error(f"Error in backward query generation: {e}")
|
|
return query
|
|
|
|
async def decompose_query(self, query: str, history: str = "") -> List[str]:
|
|
"""
|
|
Decompose complex query into simpler sub-queries
|
|
|
|
Args:
|
|
query: Original user query
|
|
history: Chat history (optional)
|
|
|
|
Returns:
|
|
List of sub-queries
|
|
"""
|
|
try:
|
|
prompt = DECOMPOSE_PROMPT.format(query=query, history=history)
|
|
|
|
response = await self.llm.acomplete(prompt=prompt)
|
|
sub_queries_text = response.text.strip()
|
|
|
|
# Parse sub-queries from response
|
|
sub_queries = []
|
|
for line in sub_queries_text.split('\n'):
|
|
line = line.strip()
|
|
if line and (line.startswith('1.') or line.startswith('2.') or line.startswith('3.') or line.startswith('4.') or line.startswith('5.')):
|
|
sub_query = line.split('.', 1)[1].strip()
|
|
if sub_query:
|
|
sub_queries.append(sub_query)
|
|
|
|
# If no valid sub-queries found, return original query as single sub-query
|
|
if not sub_queries:
|
|
sub_queries = [query]
|
|
|
|
logger.info(f"Original query: {query}")
|
|
logger.info(f"Decomposed sub-queries: {sub_queries}")
|
|
return sub_queries
|
|
except Exception as e:
|
|
logger.error(f"Error in query decomposition: {e}")
|
|
return [query]
|
|
|
|
async def transform_query(self, query: str, history: str = "") -> Dict[str, any]:
|
|
"""
|
|
Transform query using all three techniques
|
|
|
|
Args:
|
|
query: Original user query
|
|
history: Chat history (optional)
|
|
|
|
Returns:
|
|
Dict with transformed queries
|
|
"""
|
|
try:
|
|
logger.info(f"Transforming query: {query}")
|
|
|
|
# Run all transformations in parallel
|
|
from asyncio import gather
|
|
rewritten, backward, sub_queries = await gather(
|
|
self.rewrite_query(query, history),
|
|
self.backward_query(query, history),
|
|
self.decompose_query(query, history)
|
|
)
|
|
|
|
# Combine all transformed queries for comprehensive retrieval
|
|
all_transformed_queries = [rewritten, backward] + sub_queries
|
|
|
|
return {
|
|
"original": query,
|
|
"rewritten": rewritten,
|
|
"backward": backward,
|
|
"sub_queries": sub_queries,
|
|
"all_transformed": all_transformed_queries
|
|
}
|
|
except Exception as e:
|
|
logger.error(f"Error in query transformation: {e}")
|
|
return {
|
|
"original": query,
|
|
"rewritten": query,
|
|
"backward": query,
|
|
"sub_queries": [query],
|
|
"all_transformed": [query]
|
|
}
|