MrSimple07 commited on
Commit
23c53c9
·
1 Parent(s): 4606a68

hybdid ap[proach

Browse files
Files changed (1) hide show
  1. index_retriever.py +39 -16
index_retriever.py CHANGED
@@ -3,6 +3,8 @@ from llama_index.core.query_engine import RetrieverQueryEngine
3
  from llama_index.core.retrievers import VectorIndexRetriever
4
  from llama_index.core.response_synthesizers import get_response_synthesizer, ResponseMode
5
  from llama_index.core.prompts import PromptTemplate
 
 
6
  from my_logging import log_message
7
  from config import CUSTOM_PROMPT, PROMPT_SIMPLE_POISK
8
 
@@ -12,10 +14,21 @@ def create_vector_index(documents):
12
 
13
  def create_query_engine(vector_index):
14
  try:
15
- # Use only semantic/vector search with top_k=15
 
 
 
 
16
  vector_retriever = VectorIndexRetriever(
17
  index=vector_index,
18
- similarity_top_k=30
 
 
 
 
 
 
 
19
  )
20
 
21
  custom_prompt_template = PromptTemplate(PROMPT_SIMPLE_POISK)
@@ -25,39 +38,49 @@ def create_query_engine(vector_index):
25
  )
26
 
27
  query_engine = RetrieverQueryEngine(
28
- retriever=vector_retriever,
29
  response_synthesizer=response_synthesizer
30
  )
31
 
32
- log_message("Query engine успешно создан с семантическим поиском (top_k=15)")
33
  return query_engine
34
 
35
  except Exception as e:
36
  log_message(f"Ошибка создания query engine: {str(e)}")
37
  raise
38
 
39
- def rerank_nodes(query, nodes, reranker, top_k=10):
40
  if not nodes or not reranker:
41
  return nodes[:top_k]
42
 
43
  try:
44
  log_message(f"Переранжирую {len(nodes)} узлов")
45
 
46
- # Rerank ALL nodes based on relevance (including tables and images)
47
- pairs = []
48
- for node in nodes:
49
- pairs.append([query, node.text])
50
 
51
- scores = reranker.predict(pairs)
52
- scored_nodes = list(zip(nodes, scores))
53
- scored_nodes.sort(key=lambda x: x[1], reverse=True)
54
 
55
- # Return top_k most relevant nodes regardless of type
56
- result = [node for node, score in scored_nodes[:top_k]]
 
 
 
 
 
 
 
 
 
 
57
 
58
- log_message(f"Возвращаю top {len(result)} переранжированных узлов")
59
- log_message(f"Типы узлов: {[node.metadata.get('type', 'text') for node in result]}")
 
60
 
 
61
  return result
62
 
63
  except Exception as e:
 
3
  from llama_index.core.retrievers import VectorIndexRetriever
4
  from llama_index.core.response_synthesizers import get_response_synthesizer, ResponseMode
5
  from llama_index.core.prompts import PromptTemplate
6
+ from llama_index.retrievers.bm25 import BM25Retriever
7
+ from llama_index.core.retrievers import QueryFusionRetriever
8
  from my_logging import log_message
9
  from config import CUSTOM_PROMPT, PROMPT_SIMPLE_POISK
10
 
 
14
 
15
  def create_query_engine(vector_index):
16
  try:
17
+ bm25_retriever = BM25Retriever.from_defaults(
18
+ docstore=vector_index.docstore,
19
+ similarity_top_k=15
20
+ )
21
+
22
  vector_retriever = VectorIndexRetriever(
23
  index=vector_index,
24
+ similarity_top_k=30,
25
+ similarity_cutoff=0.8
26
+ )
27
+
28
+ hybrid_retriever = QueryFusionRetriever(
29
+ [vector_retriever, bm25_retriever],
30
+ similarity_top_k=30,
31
+ num_queries=1
32
  )
33
 
34
  custom_prompt_template = PromptTemplate(PROMPT_SIMPLE_POISK)
 
38
  )
39
 
40
  query_engine = RetrieverQueryEngine(
41
+ retriever=hybrid_retriever,
42
  response_synthesizer=response_synthesizer
43
  )
44
 
45
+ log_message("Query engine успешно создан")
46
  return query_engine
47
 
48
  except Exception as e:
49
  log_message(f"Ошибка создания query engine: {str(e)}")
50
  raise
51
 
52
+ def rerank_nodes(query, nodes, reranker, top_k=15):
53
  if not nodes or not reranker:
54
  return nodes[:top_k]
55
 
56
  try:
57
  log_message(f"Переранжирую {len(nodes)} узлов")
58
 
59
+ # Separate tables and images from text nodes
60
+ table_nodes = [node for node in nodes if node.metadata.get('type') == 'table']
61
+ image_nodes = [node for node in nodes if node.metadata.get('type') == 'image']
62
+ text_nodes = [node for node in nodes if node.metadata.get('type', 'text') == 'text']
63
 
64
+ priority_nodes = table_nodes + image_nodes
 
 
65
 
66
+ # Rerank only text nodes
67
+ if text_nodes:
68
+ pairs = []
69
+ for node in text_nodes:
70
+ pairs.append([query, node.text])
71
+
72
+ scores = reranker.predict(pairs)
73
+ scored_nodes = list(zip(text_nodes, scores))
74
+ scored_nodes.sort(key=lambda x: x[1], reverse=True)
75
+ reranked_text_nodes = [node for node, score in scored_nodes]
76
+ else:
77
+ reranked_text_nodes = []
78
 
79
+ # Combine: priority nodes first, then reranked text nodes
80
+ final_nodes = priority_nodes + reranked_text_nodes
81
+ result = final_nodes[:top_k]
82
 
83
+ log_message(f"Возвращаю {len(priority_nodes)} приоритетных узлов и {len(result) - len(priority_nodes)} текстовых узлов")
84
  return result
85
 
86
  except Exception as e: