Spaces:
Sleeping
Sleeping
Commit ·
67e6e17
1
Parent(s): 63590e6
hybrid search - Qdrant + BM25 + RRF, switch LLM to Groq
Browse files- src/rag.py +9 -6
- src/vector_store.py +99 -35
src/rag.py
CHANGED
|
@@ -1,13 +1,13 @@
|
|
| 1 |
import os
|
| 2 |
-
from
|
| 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 =
|
| 9 |
|
| 10 |
-
LLM_MODEL = "
|
| 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.
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
return {
|
| 39 |
-
"answer": response.
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
from tqdm import tqdm
|
| 4 |
from src.embedder import embed_texts
|
| 5 |
|
| 6 |
-
|
| 7 |
COLLECTION_NAME = "rag_docs"
|
|
|
|
| 8 |
BATCH_SIZE = 20
|
|
|
|
| 9 |
|
| 10 |
|
| 11 |
-
def
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
|
| 18 |
|
| 19 |
def add_chunks(chunks: list[dict]):
|
| 20 |
-
"""Embed và lưu chunks vào
|
| 21 |
-
|
|
|
|
|
|
|
| 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
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 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
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|