Spaces:
Sleeping
Sleeping
Commit ·
d287ce2
1
Parent(s): c696a51
Vietnamese tokenization (underthesea) + query rewriting
Browse files- ingest.py +9 -6
- src/query_rewriter.py +27 -0
- src/rag.py +8 -4
- src/vector_store.py +11 -2
ingest.py
CHANGED
|
@@ -6,14 +6,17 @@ IMAGE_FOLDER = "data/images"
|
|
| 6 |
if __name__ == "__main__":
|
| 7 |
mode = sys.argv[1] if len(sys.argv) > 1 else "all"
|
| 8 |
|
| 9 |
-
if mode
|
| 10 |
-
import
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
| 12 |
from src.parsers.router import parse_folder as parse_folder_pdf
|
| 13 |
from src.chunker import chunk_pages
|
| 14 |
from src.vector_store import add_chunks
|
| 15 |
|
| 16 |
-
print("=== INGEST PDF (
|
| 17 |
pdf_pages = parse_folder_pdf(PDF_FOLDER, extensions={".pdf"})
|
| 18 |
print(f"Tổng: {len(pdf_pages)} đoạn từ PDF")
|
| 19 |
chunks = chunk_pages(pdf_pages)
|
|
@@ -22,7 +25,7 @@ if __name__ == "__main__":
|
|
| 22 |
add_chunks(chunks)
|
| 23 |
print("PDF xong!\n")
|
| 24 |
|
| 25 |
-
|
| 26 |
from src.parsers.router import parse_folder as parse_folder_img
|
| 27 |
from src.chunker import chunk_pages
|
| 28 |
from src.vector_store import add_chunks
|
|
@@ -34,4 +37,4 @@ if __name__ == "__main__":
|
|
| 34 |
print(f"Tổng: {len(chunks)} chunks")
|
| 35 |
if chunks:
|
| 36 |
add_chunks(chunks)
|
| 37 |
-
print("Ảnh xong!")
|
|
|
|
| 6 |
if __name__ == "__main__":
|
| 7 |
mode = sys.argv[1] if len(sys.argv) > 1 else "all"
|
| 8 |
|
| 9 |
+
if mode == "all":
|
| 10 |
+
import subprocess
|
| 11 |
+
subprocess.run([sys.executable, __file__, "pdf"], check=True)
|
| 12 |
+
subprocess.run([sys.executable, __file__, "images"], check=True)
|
| 13 |
+
|
| 14 |
+
elif mode == "pdf":
|
| 15 |
from src.parsers.router import parse_folder as parse_folder_pdf
|
| 16 |
from src.chunker import chunk_pages
|
| 17 |
from src.vector_store import add_chunks
|
| 18 |
|
| 19 |
+
print("=== INGEST PDF (pypdf - CPU) ===")
|
| 20 |
pdf_pages = parse_folder_pdf(PDF_FOLDER, extensions={".pdf"})
|
| 21 |
print(f"Tổng: {len(pdf_pages)} đoạn từ PDF")
|
| 22 |
chunks = chunk_pages(pdf_pages)
|
|
|
|
| 25 |
add_chunks(chunks)
|
| 26 |
print("PDF xong!\n")
|
| 27 |
|
| 28 |
+
elif mode == "images":
|
| 29 |
from src.parsers.router import parse_folder as parse_folder_img
|
| 30 |
from src.chunker import chunk_pages
|
| 31 |
from src.vector_store import add_chunks
|
|
|
|
| 37 |
print(f"Tổng: {len(chunks)} chunks")
|
| 38 |
if chunks:
|
| 39 |
add_chunks(chunks)
|
| 40 |
+
print("Ảnh xong!")
|
src/query_rewriter.py
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from groq import Groq
|
| 2 |
+
import os
|
| 3 |
+
from dotenv import load_dotenv
|
| 4 |
+
|
| 5 |
+
load_dotenv()
|
| 6 |
+
_client = Groq(api_key=os.environ["GROQ_API_KEY"])
|
| 7 |
+
|
| 8 |
+
REWRITE_MODEL = "llama-3.3-70b-versatile"
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def rewrite_query(query: str) -> str:
|
| 12 |
+
"""Paraphrase câu hỏi để tăng khả năng tìm kiếm."""
|
| 13 |
+
prompt = f"""Bạn là công cụ cải thiện câu truy vấn tìm kiếm tiếng Việt.
|
| 14 |
+
Hãy viết lại câu hỏi sau thành một câu truy vấn ngắn gọn, súc tích hơn với các từ khóa quan trọng.
|
| 15 |
+
Chỉ trả về câu truy vấn mới, không giải thích.
|
| 16 |
+
|
| 17 |
+
Câu hỏi gốc: {query}
|
| 18 |
+
Câu truy vấn mới:"""
|
| 19 |
+
|
| 20 |
+
response = _client.chat.completions.create(
|
| 21 |
+
model=REWRITE_MODEL,
|
| 22 |
+
messages=[{"role": "user", "content": prompt}],
|
| 23 |
+
max_tokens=100,
|
| 24 |
+
temperature=0.3,
|
| 25 |
+
)
|
| 26 |
+
rewritten = response.choices[0].message.content.strip()
|
| 27 |
+
return rewritten if rewritten else query
|
src/rag.py
CHANGED
|
@@ -4,6 +4,7 @@ from dotenv import load_dotenv
|
|
| 4 |
from src.embedder import embed_query
|
| 5 |
from src.vector_store import query as vector_query
|
| 6 |
from src.reranker import rerank
|
|
|
|
| 7 |
|
| 8 |
load_dotenv()
|
| 9 |
client = Groq(api_key=os.environ["GROQ_API_KEY"])
|
|
@@ -30,10 +31,13 @@ TRẢ LỜI:"""
|
|
| 30 |
|
| 31 |
|
| 32 |
def answer(question: str, top_k: int = 5) -> dict:
|
| 33 |
-
"""RAG pipeline: embed -> hybrid retrieve -> rerank -> generate."""
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
|
|
|
|
|
|
|
|
|
| 37 |
contexts = rerank(question, candidates, top_k=top_k)
|
| 38 |
|
| 39 |
prompt = build_prompt(question, contexts)
|
|
|
|
| 4 |
from src.embedder import embed_query
|
| 5 |
from src.vector_store import query as vector_query
|
| 6 |
from src.reranker import rerank
|
| 7 |
+
from src.query_rewriter import rewrite_query
|
| 8 |
|
| 9 |
load_dotenv()
|
| 10 |
client = Groq(api_key=os.environ["GROQ_API_KEY"])
|
|
|
|
| 31 |
|
| 32 |
|
| 33 |
def answer(question: str, top_k: int = 5) -> dict:
|
| 34 |
+
"""RAG pipeline: rewrite -> embed -> hybrid retrieve -> rerank -> generate."""
|
| 35 |
+
rewritten = rewrite_query(question)
|
| 36 |
+
if rewritten != question:
|
| 37 |
+
print(f"Query rewritten: {rewritten}")
|
| 38 |
+
|
| 39 |
+
query_vec = embed_query(rewritten)
|
| 40 |
+
candidates = vector_query(query_vec, query_text=rewritten, top_k=top_k * 3)
|
| 41 |
contexts = rerank(question, candidates, top_k=top_k)
|
| 42 |
|
| 43 |
prompt = build_prompt(question, contexts)
|
src/vector_store.py
CHANGED
|
@@ -5,8 +5,17 @@ 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"
|
|
@@ -53,7 +62,7 @@ def add_chunks(chunks: list[dict]):
|
|
| 53 |
|
| 54 |
# Build BM25 index
|
| 55 |
print("Building BM25 index...")
|
| 56 |
-
tokenized = [
|
| 57 |
bm25 = BM25Okapi(tokenized)
|
| 58 |
with open(BM25_PATH, "wb") as f:
|
| 59 |
pickle.dump({"bm25": bm25, "texts": texts, "metadatas": metadatas}, f)
|
|
@@ -95,7 +104,7 @@ def query(query_embedding: list[float], query_text: str, top_k: int = 5) -> list
|
|
| 95 |
texts = bm25_data["texts"]
|
| 96 |
metadatas = bm25_data["metadatas"]
|
| 97 |
|
| 98 |
-
tokenized_query =
|
| 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 |
|
|
|
|
| 5 |
from qdrant_client.models import Distance, VectorParams, PointStruct
|
| 6 |
from rank_bm25 import BM25Okapi
|
| 7 |
from tqdm import tqdm
|
| 8 |
+
from underthesea import word_tokenize
|
| 9 |
from src.embedder import embed_texts
|
| 10 |
|
| 11 |
+
|
| 12 |
+
def _tokenize_vi(text: str) -> list[str]:
|
| 13 |
+
"""Tách từ tiếng Việt, fallback về split nếu lỗi."""
|
| 14 |
+
try:
|
| 15 |
+
return word_tokenize(text.lower(), format="text").split()
|
| 16 |
+
except Exception:
|
| 17 |
+
return text.lower().split()
|
| 18 |
+
|
| 19 |
QDRANT_PATH = "qdrant_storage"
|
| 20 |
COLLECTION_NAME = "rag_docs"
|
| 21 |
BM25_PATH = "bm25_index.pkl"
|
|
|
|
| 62 |
|
| 63 |
# Build BM25 index
|
| 64 |
print("Building BM25 index...")
|
| 65 |
+
tokenized = [_tokenize_vi(t) for t in texts]
|
| 66 |
bm25 = BM25Okapi(tokenized)
|
| 67 |
with open(BM25_PATH, "wb") as f:
|
| 68 |
pickle.dump({"bm25": bm25, "texts": texts, "metadatas": metadatas}, f)
|
|
|
|
| 104 |
texts = bm25_data["texts"]
|
| 105 |
metadatas = bm25_data["metadatas"]
|
| 106 |
|
| 107 |
+
tokenized_query = _tokenize_vi(query_text)
|
| 108 |
bm25_scores = bm25.get_scores(tokenized_query)
|
| 109 |
bm25_ids = sorted(range(len(bm25_scores)), key=lambda i: bm25_scores[i], reverse=True)[:top_n]
|
| 110 |
|