thaidinhz1 commited on
Commit
d287ce2
·
1 Parent(s): c696a51

Vietnamese tokenization (underthesea) + query rewriting

Browse files
Files changed (4) hide show
  1. ingest.py +9 -6
  2. src/query_rewriter.py +27 -0
  3. src/rag.py +8 -4
  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 in ("pdf", "all"):
10
- import os
11
- os.environ["CUDA_VISIBLE_DEVICES"] = "" # force CPU cho Docling
 
 
 
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 (Docling - CPU) ===")
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
- if mode in ("images", "all"):
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
- query_vec = embed_query(question)
35
- # Lấy nhiều hơn để reranker có đủ candidates
36
- candidates = vector_query(query_vec, query_text=question, top_k=top_k * 3)
 
 
 
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 = [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)
@@ -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 = 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
 
 
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