Create rag.py
Browse files
rag.py
CHANGED
|
@@ -1,37 +1,55 @@
|
|
| 1 |
-
|
| 2 |
import faiss
|
|
|
|
| 3 |
from transformers import AutoTokenizer, AutoModelForCausalLM
|
| 4 |
-
import torch
|
| 5 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 6 |
|
| 7 |
embed_model = SentenceTransformer("BAAI/bge-small-zh-v1.5")
|
| 8 |
|
|
|
|
|
|
|
|
|
|
| 9 |
|
| 10 |
-
|
|
|
|
| 11 |
|
|
|
|
|
|
|
| 12 |
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
blocks = raw.split("===")
|
| 16 |
|
|
|
|
|
|
|
| 17 |
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
|
|
|
|
|
|
|
| 22 |
|
| 23 |
-
|
|
|
|
|
|
|
| 24 |
|
|
|
|
|
|
|
|
|
|
| 25 |
|
| 26 |
-
|
| 27 |
-
q_emb = embed_model.encode([question], normalize_embeddings=True)
|
| 28 |
-
_, I = index.search(q_emb, 1)
|
| 29 |
-
context = blocks[I[0][0]]
|
| 30 |
-
|
| 31 |
|
| 32 |
-
|
|
|
|
|
|
|
|
|
|
| 33 |
|
|
|
|
| 34 |
|
| 35 |
-
inputs = tokenizer(prompt, return_tensors="pt")
|
| 36 |
-
outputs = model.generate(**inputs, max_new_tokens=200)
|
| 37 |
-
return tokenizer.decode(outputs[0], skip_special_tokens=True)
|
|
|
|
| 1 |
+
import os
|
| 2 |
import faiss
|
| 3 |
+
from sentence_transformers import SentenceTransformer
|
| 4 |
from transformers import AutoTokenizer, AutoModelForCausalLM
|
|
|
|
| 5 |
|
| 6 |
+
DATA_PATH = "data/feng_family.txt"
|
| 7 |
+
INDEX_DIR = "index"
|
| 8 |
+
FAISS_PATH = os.path.join(INDEX_DIR, "family.faiss")
|
| 9 |
+
BLOCKS_PATH = os.path.join(INDEX_DIR, "blocks.txt")
|
| 10 |
|
| 11 |
embed_model = SentenceTransformer("BAAI/bge-small-zh-v1.5")
|
| 12 |
|
| 13 |
+
def build_index_if_needed():
|
| 14 |
+
if os.path.exists(FAISS_PATH) and os.path.exists(BLOCKS_PATH):
|
| 15 |
+
return
|
| 16 |
|
| 17 |
+
with open(DATA_PATH, "r", encoding="utf-8") as f:
|
| 18 |
+
text = f.read()
|
| 19 |
|
| 20 |
+
blocks = [b.strip() for b in text.split("--- PERSON ---") if b.strip()]
|
| 21 |
+
embeddings = embed_model.encode(blocks, normalize_embeddings=True)
|
| 22 |
|
| 23 |
+
index = faiss.IndexFlatIP(embeddings.shape[1])
|
| 24 |
+
index.add(embeddings)
|
|
|
|
| 25 |
|
| 26 |
+
os.makedirs(INDEX_DIR, exist_ok=True)
|
| 27 |
+
faiss.write_index(index, FAISS_PATH)
|
| 28 |
|
| 29 |
+
with open(BLOCKS_PATH, "w", encoding="utf-8") as f:
|
| 30 |
+
for b in blocks:
|
| 31 |
+
f.write(b.replace("\n", " ") + "\n===\n")
|
| 32 |
|
| 33 |
+
# 启动时自动建索引(只会跑一次)
|
| 34 |
+
build_index_if_needed()
|
| 35 |
|
| 36 |
+
index = faiss.read_index(FAISS_PATH)
|
| 37 |
+
with open(BLOCKS_PATH, "r", encoding="utf-8") as f:
|
| 38 |
+
blocks = f.read().split("===")
|
| 39 |
|
| 40 |
+
llm_name = "Qwen/Qwen2.5-1.5B"
|
| 41 |
+
tokenizer = AutoTokenizer.from_pretrained(llm_name)
|
| 42 |
+
model = AutoModelForCausalLM.from_pretrained(llm_name)
|
| 43 |
|
| 44 |
+
SYSTEM_PROMPT = "你是冯氏家谱机器人,只能根据家谱资料回答,没有记载就说未记载。"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
|
| 46 |
+
def ask(question: str) -> str:
|
| 47 |
+
q_emb = embed_model.encode([question], normalize_embeddings=True)
|
| 48 |
+
_, I = index.search(q_emb, 1)
|
| 49 |
+
context = blocks[I[0][0]]
|
| 50 |
|
| 51 |
+
prompt = f"{SYSTEM_PROMPT}\n\n【家谱资料】\n{context}\n\n【问题】{question}\n回答:"
|
| 52 |
|
| 53 |
+
inputs = tokenizer(prompt, return_tensors="pt")
|
| 54 |
+
outputs = model.generate(**inputs, max_new_tokens=200)
|
| 55 |
+
return tokenizer.decode(outputs[0], skip_special_tokens=True)
|