Spaces:
Sleeping
Sleeping
Download src/vector_store.py from thaidinhz1/rag-vietnamese: direct link, hf CLI and curl.
- Browser
- Download file 4.42 kB
-
https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/vector_store.py
- Command line
-
hf download hf://spaces/thaidinhz1/rag-vietnamese@af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/vector_store.py
-
curl -L -o vector_store.py https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/vector_store.py
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 | |