rag-vietnamese / src /vector_store.py
thaidinhz1's picture
fix: commit bm25_index.pkl for HF Spaces deployment
a088361
Raw History Blame
4.42 kB
import os
import time
import pickle
from pathlib import Path
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, VectorParams, PointStruct
from rank_bm25 import BM25Okapi
from tqdm import tqdm
from underthesea import word_tokenize
from src.embedder import embed_texts
def _tokenize_vi(text: str) -> list[str]:
"""Tách từ tiếng Việt, fallback về split nếu lỗi."""
try:
return word_tokenize(text.lower(), format="text").split()
except Exception:
return text.lower().split()
QDRANT_PATH = "qdrant_storage"
COLLECTION_NAME = "rag_docs"
BM25_PATH = "bm25_index.pkl"
BATCH_SIZE = 20
VECTOR_DIM = 3072 # gemini-embedding-001
def get_client():
url = os.getenv("QDRANT_URL")
api_key = os.getenv("QDRANT_API_KEY")
if url and api_key:
return QdrantClient(url=url, api_key=api_key, timeout=60)
return QdrantClient(path=QDRANT_PATH)
def get_or_create_collection(client):
existing = [c.name for c in client.get_collections().collections]
if COLLECTION_NAME not in existing:
client.create_collection(
collection_name=COLLECTION_NAME,
vectors_config=VectorParams(size=VECTOR_DIM, distance=Distance.COSINE),
)
return client.get_collection(COLLECTION_NAME)
def add_chunks(chunks: list[dict]):
"""Embed và lưu chunks vào Qdrant + build BM25 index."""
client = get_client()
get_or_create_collection(client)
texts = [c["text"] for c in chunks]
metadatas = [c["metadata"] for c in chunks]
print(f"Embedding {len(chunks)} chunks...")
all_embeddings = []
for i in tqdm(range(0, len(chunks), BATCH_SIZE)):
batch_texts = texts[i:i + BATCH_SIZE]
embeddings = embed_texts(batch_texts)
all_embeddings.extend(embeddings)
time.sleep(BATCH_SIZE * 1.5)
# Lưu vào Qdrant
points = [
PointStruct(id=i, vector=emb, payload={"text": texts[i], **metadatas[i]})
for i, emb in enumerate(all_embeddings)
]
client.upsert(collection_name=COLLECTION_NAME, points=points)
# Build BM25 index
print("Building BM25 index...")
tokenized = [_tokenize_vi(t) for t in texts]
bm25 = BM25Okapi(tokenized)
with open(BM25_PATH, "wb") as f:
pickle.dump({"bm25": bm25, "texts": texts, "metadatas": metadatas}, f)
print(f"Đã lưu {len(chunks)} chunks vào Qdrant + BM25")
def _load_bm25():
with open(BM25_PATH, "rb") as f:
return pickle.load(f)
def _reciprocal_rank_fusion(rankings: list[list[int]], k: int = 60) -> list[tuple[int, float]]:
"""Kết hợp nhiều bảng xếp hạng bằng RRF."""
scores = {}
for ranking in rankings:
for rank, doc_id in enumerate(ranking):
scores[doc_id] = scores.get(doc_id, 0) + 1 / (k + rank + 1)
return sorted(scores.items(), key=lambda x: x[1], reverse=True)
def query(query_embedding: list[float], query_text: str, top_k: int = 5) -> list[dict]:
"""Hybrid retrieval: vector + BM25 + RRF."""
client = get_client()
top_n = top_k * 4 # lấy nhiều hơn để RRF có đủ candidates
# Vector search
vector_results = client.query_points(
collection_name=COLLECTION_NAME,
query=query_embedding,
limit=top_n,
).points
vector_ids = [r.id for r in vector_results]
vector_payload = {r.id: r.payload for r in vector_results}
# BM25 search
bm25_data = _load_bm25()
bm25 = bm25_data["bm25"]
texts = bm25_data["texts"]
metadatas = bm25_data["metadatas"]
tokenized_query = _tokenize_vi(query_text)
bm25_scores = bm25.get_scores(tokenized_query)
bm25_ids = sorted(range(len(bm25_scores)), key=lambda i: bm25_scores[i], reverse=True)[:top_n]
# RRF
fused = _reciprocal_rank_fusion([vector_ids, bm25_ids])[:top_k]
# Build kết quả
hits = []
for doc_id, rrf_score in fused:
if doc_id in vector_payload:
payload = vector_payload[doc_id]
hits.append({
"text": payload["text"],
"metadata": {k: v for k, v in payload.items() if k != "text"},
"score": round(rrf_score, 4),
})
elif doc_id < len(texts):
hits.append({
"text": texts[doc_id],
"metadata": metadatas[doc_id],
"score": round(rrf_score, 4),
})
return hits