Spaces:
Sleeping
Sleeping
Download src/rag.py from thaidinhz1/rag-vietnamese: direct link, hf CLI and curl.
- Browser
- Download file 2.41 kB
-
https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/rag.py
- Command line
-
hf download hf://spaces/thaidinhz1/rag-vietnamese@af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/rag.py
-
curl -L -o rag.py https://huggingface.co/spaces/thaidinhz1/rag-vietnamese/resolve/af92b4d4054d18f33dedfaeacdbb504078b2e72c/src/rag.py
2.41 kB
| import os | |
| from groq import Groq | |
| from dotenv import load_dotenv | |
| from src.embedder import embed_query | |
| from src.vector_store import query as vector_query | |
| from src.reranker import rerank | |
| from src.query_rewriter import rewrite_query | |
| from src.uploader import search_uploaded | |
| load_dotenv() | |
| client = Groq(api_key=os.environ["GROQ_API_KEY"].strip()) | |
| LLM_MODEL = "llama-3.3-70b-versatile" | |
| def build_prompt(question: str, contexts: list[dict]) -> str: | |
| context_text = "" | |
| for i, ctx in enumerate(contexts, 1): | |
| meta = ctx["metadata"] | |
| context_text += f"[{i}] (File: {meta['source']}, Trang: {meta['page']})\n{ctx['text']}\n\n" | |
| return f"""Bạn là trợ lý trả lời câu hỏi dựa trên tài liệu được cung cấp. | |
| Chỉ dùng thông tin trong phần NGỮ CẢNH bên dưới để trả lời. | |
| Trích dẫn nguồn bằng số [1], [2],... tương ứng với từng đoạn bạn sử dụng. | |
| Nếu không tìm thấy thông tin, hãy nói "Không tìm thấy thông tin trong tài liệu." | |
| NGỮ CẢNH: | |
| {context_text} | |
| CÂU HỎI: {question} | |
| TRẢ LỜI:""" | |
| def retrieve(question: str, top_k: int = 5): | |
| """Rewrite -> embed -> hybrid retrieve -> rerank. Returns (contexts, sources).""" | |
| rewritten = rewrite_query(question) | |
| if rewritten != question: | |
| print(f"Query rewritten: {rewritten}") | |
| query_vec = embed_query(rewritten) | |
| candidates = vector_query(query_vec, query_text=rewritten, top_k=top_k * 3) | |
| uploaded = search_uploaded(query_vec, rewritten, top_k=top_k * 2) | |
| contexts = rerank(question, candidates + uploaded, top_k=top_k) | |
| sources = [ | |
| { | |
| "source": c["metadata"]["source"], | |
| "page": c["metadata"]["page"], | |
| "rrf_score": round(c["score"], 3), | |
| "rerank_score": round(c["rerank_score"], 3), | |
| } | |
| for c in contexts | |
| ] | |
| return contexts, sources | |
| def answer(question: str, top_k: int = 5) -> dict: | |
| """RAG pipeline: rewrite -> embed -> hybrid retrieve -> rerank -> generate.""" | |
| contexts, sources = retrieve(question, top_k) | |
| prompt = build_prompt(question, contexts) | |
| response = client.chat.completions.create( | |
| model=LLM_MODEL, | |
| messages=[{"role": "user", "content": prompt}], | |
| ) | |
| return { | |
| "answer": response.choices[0].message.content, | |
| "contexts": [c["text"] for c in contexts], | |
| "sources": sources, | |
| } |