RAG/utils/query_transformer.py

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]
}