thaidinhz1 commited on
Commit
67e6e17
·
1 Parent(s): 63590e6

hybrid search - Qdrant + BM25 + RRF, switch LLM to Groq

Browse files
Files changed (2) hide show
  1. src/rag.py +9 -6
  2. src/vector_store.py +99 -35
src/rag.py CHANGED
@@ -1,13 +1,13 @@
1
  import os
2
- from google import genai
3
  from dotenv import load_dotenv
4
  from src.embedder import embed_query
5
  from src.vector_store import query as vector_query
6
 
7
  load_dotenv()
8
- client = genai.Client(api_key=os.environ["GOOGLE_API_KEY"])
9
 
10
- LLM_MODEL = "gemini-2.5-flash"
11
 
12
 
13
  def build_prompt(question: str, contexts: list[dict]) -> str:
@@ -30,13 +30,16 @@ TRẢ LỜI:"""
30
  def answer(question: str, top_k: int = 5) -> dict:
31
  """RAG pipeline: embed query -> retrieve -> generate."""
32
  query_vec = embed_query(question)
33
- contexts = vector_query(query_vec, top_k=top_k)
34
 
35
  prompt = build_prompt(question, contexts)
36
- response = client.models.generate_content(model=LLM_MODEL, contents=prompt)
 
 
 
37
 
38
  return {
39
- "answer": response.text,
40
  "sources": [
41
  {"source": c["metadata"]["source"], "page": c["metadata"]["page"], "score": round(c["score"], 3)}
42
  for c in contexts
 
1
  import os
2
+ from groq import Groq
3
  from dotenv import load_dotenv
4
  from src.embedder import embed_query
5
  from src.vector_store import query as vector_query
6
 
7
  load_dotenv()
8
+ client = Groq(api_key=os.environ["GROQ_API_KEY"])
9
 
10
+ LLM_MODEL = "llama-3.3-70b-versatile"
11
 
12
 
13
  def build_prompt(question: str, contexts: list[dict]) -> str:
 
30
  def answer(question: str, top_k: int = 5) -> dict:
31
  """RAG pipeline: embed query -> retrieve -> generate."""
32
  query_vec = embed_query(question)
33
+ contexts = vector_query(query_vec, query_text=question, top_k=top_k)
34
 
35
  prompt = build_prompt(question, contexts)
36
+ response = client.chat.completions.create(
37
+ model=LLM_MODEL,
38
+ messages=[{"role": "user", "content": prompt}],
39
+ )
40
 
41
  return {
42
+ "answer": response.choices[0].message.content,
43
  "sources": [
44
  {"source": c["metadata"]["source"], "page": c["metadata"]["page"], "score": round(c["score"], 3)}
45
  for c in contexts
src/vector_store.py CHANGED
@@ -1,57 +1,121 @@
1
  import time
2
- import chromadb
 
 
 
 
3
  from tqdm import tqdm
4
  from src.embedder import embed_texts
5
 
6
- CHROMA_PATH = "chroma_db"
7
  COLLECTION_NAME = "rag_docs"
 
8
  BATCH_SIZE = 20
 
9
 
10
 
11
- def get_collection():
12
- client = chromadb.PersistentClient(path=CHROMA_PATH)
13
- return client.get_or_create_collection(
14
- name=COLLECTION_NAME,
15
- metadata={"hnsw:space": "cosine"},
16
- )
 
 
 
 
 
 
17
 
18
 
19
  def add_chunks(chunks: list[dict]):
20
- """Embed và lưu chunks vào Chroma."""
21
- collection = get_collection()
 
 
22
  texts = [c["text"] for c in chunks]
23
  metadatas = [c["metadata"] for c in chunks]
24
- ids = [f"{m['source']}_p{m['page']}_c{m['chunk_index']}" for m in metadatas]
25
 
26
- print(f"Embedding {len(chunks)} chunks (rate limit: ~80 req/min)...")
 
27
  for i in tqdm(range(0, len(chunks), BATCH_SIZE)):
28
  batch_texts = texts[i:i + BATCH_SIZE]
29
- batch_meta = metadatas[i:i + BATCH_SIZE]
30
- batch_ids = ids[i:i + BATCH_SIZE]
31
  embeddings = embed_texts(batch_texts)
32
- collection.add(
33
- documents=batch_texts,
34
- embeddings=embeddings,
35
- metadatas=batch_meta,
36
- ids=batch_ids,
37
- )
38
- # Free tier: 100 req/min, mỗi batch=BATCH_SIZE req → nghỉ để tránh vượt quota
39
  time.sleep(BATCH_SIZE * 1.5)
40
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41
 
42
- def query(query_embedding: list[float], top_k: int = 5) -> list[dict]:
43
- """Tìm top-k chunks gần nhất."""
44
- collection = get_collection()
45
- results = collection.query(
46
- query_embeddings=[query_embedding],
47
- n_results=top_k,
48
- include=["documents", "metadatas", "distances"],
49
- )
50
  hits = []
51
- for doc, meta, dist in zip(
52
- results["documents"][0],
53
- results["metadatas"][0],
54
- results["distances"][0],
55
- ):
56
- hits.append({"text": doc, "metadata": meta, "score": 1 - dist})
 
 
 
 
 
 
 
 
57
  return hits
 
1
  import time
2
+ import pickle
3
+ from pathlib import Path
4
+ from qdrant_client import QdrantClient
5
+ from qdrant_client.models import Distance, VectorParams, PointStruct
6
+ from rank_bm25 import BM25Okapi
7
  from tqdm import tqdm
8
  from src.embedder import embed_texts
9
 
10
+ QDRANT_PATH = "qdrant_storage"
11
  COLLECTION_NAME = "rag_docs"
12
+ BM25_PATH = "bm25_index.pkl"
13
  BATCH_SIZE = 20
14
+ VECTOR_DIM = 3072 # gemini-embedding-001
15
 
16
 
17
+ def get_client():
18
+ return QdrantClient(path=QDRANT_PATH)
19
+
20
+
21
+ def get_or_create_collection(client):
22
+ existing = [c.name for c in client.get_collections().collections]
23
+ if COLLECTION_NAME not in existing:
24
+ client.create_collection(
25
+ collection_name=COLLECTION_NAME,
26
+ vectors_config=VectorParams(size=VECTOR_DIM, distance=Distance.COSINE),
27
+ )
28
+ return client.get_collection(COLLECTION_NAME)
29
 
30
 
31
  def add_chunks(chunks: list[dict]):
32
+ """Embed và lưu chunks vào Qdrant + build BM25 index."""
33
+ client = get_client()
34
+ get_or_create_collection(client)
35
+
36
  texts = [c["text"] for c in chunks]
37
  metadatas = [c["metadata"] for c in chunks]
 
38
 
39
+ print(f"Embedding {len(chunks)} chunks...")
40
+ all_embeddings = []
41
  for i in tqdm(range(0, len(chunks), BATCH_SIZE)):
42
  batch_texts = texts[i:i + BATCH_SIZE]
 
 
43
  embeddings = embed_texts(batch_texts)
44
+ all_embeddings.extend(embeddings)
 
 
 
 
 
 
45
  time.sleep(BATCH_SIZE * 1.5)
46
 
47
+ # Lưu vào Qdrant
48
+ points = [
49
+ PointStruct(id=i, vector=emb, payload={"text": texts[i], **metadatas[i]})
50
+ for i, emb in enumerate(all_embeddings)
51
+ ]
52
+ client.upsert(collection_name=COLLECTION_NAME, points=points)
53
+
54
+ # Build BM25 index
55
+ print("Building BM25 index...")
56
+ tokenized = [t.lower().split() for t in texts]
57
+ bm25 = BM25Okapi(tokenized)
58
+ with open(BM25_PATH, "wb") as f:
59
+ pickle.dump({"bm25": bm25, "texts": texts, "metadatas": metadatas}, f)
60
+
61
+ print(f"Đã lưu {len(chunks)} chunks vào Qdrant + BM25")
62
+
63
+
64
+ def _load_bm25():
65
+ with open(BM25_PATH, "rb") as f:
66
+ return pickle.load(f)
67
+
68
+
69
+ def _reciprocal_rank_fusion(rankings: list[list[int]], k: int = 60) -> list[tuple[int, float]]:
70
+ """Kết hợp nhiều bảng xếp hạng bằng RRF."""
71
+ scores = {}
72
+ for ranking in rankings:
73
+ for rank, doc_id in enumerate(ranking):
74
+ scores[doc_id] = scores.get(doc_id, 0) + 1 / (k + rank + 1)
75
+ return sorted(scores.items(), key=lambda x: x[1], reverse=True)
76
+
77
+
78
+ def query(query_embedding: list[float], query_text: str, top_k: int = 5) -> list[dict]:
79
+ """Hybrid retrieval: vector + BM25 + RRF."""
80
+ client = get_client()
81
+ top_n = top_k * 4 # lấy nhiều hơn để RRF có đủ candidates
82
+
83
+ # Vector search
84
+ vector_results = client.query_points(
85
+ collection_name=COLLECTION_NAME,
86
+ query=query_embedding,
87
+ limit=top_n,
88
+ ).points
89
+ vector_ids = [r.id for r in vector_results]
90
+ vector_payload = {r.id: r.payload for r in vector_results}
91
+
92
+ # BM25 search
93
+ bm25_data = _load_bm25()
94
+ bm25 = bm25_data["bm25"]
95
+ texts = bm25_data["texts"]
96
+ metadatas = bm25_data["metadatas"]
97
+
98
+ tokenized_query = query_text.lower().split()
99
+ bm25_scores = bm25.get_scores(tokenized_query)
100
+ bm25_ids = sorted(range(len(bm25_scores)), key=lambda i: bm25_scores[i], reverse=True)[:top_n]
101
+
102
+ # RRF
103
+ fused = _reciprocal_rank_fusion([vector_ids, bm25_ids])[:top_k]
104
 
105
+ # Build kết quả
 
 
 
 
 
 
 
106
  hits = []
107
+ for doc_id, rrf_score in fused:
108
+ if doc_id in vector_payload:
109
+ payload = vector_payload[doc_id]
110
+ hits.append({
111
+ "text": payload["text"],
112
+ "metadata": {k: v for k, v in payload.items() if k != "text"},
113
+ "score": round(rrf_score, 4),
114
+ })
115
+ elif doc_id < len(texts):
116
+ hits.append({
117
+ "text": texts[doc_id],
118
+ "metadata": metadatas[doc_id],
119
+ "score": round(rrf_score, 4),
120
+ })
121
  return hits