fdbw commited on
Commit
9df46ad
·
verified ·
1 Parent(s): 154d697

Create rag.py

Browse files
Files changed (1) hide show
  1. rag.py +37 -19
rag.py CHANGED
@@ -1,37 +1,55 @@
1
- from sentence_transformers import SentenceTransformer
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
- index = faiss.read_index("index/family.faiss")
 
11
 
 
 
12
 
13
- with open("index/blocks.txt", "r", encoding="utf-8") as f:
14
- raw = f.read()
15
- blocks = raw.split("===")
16
 
 
 
17
 
18
- llm_name = "Qwen/Qwen2.5-1.5B"
19
- tokenizer = AutoTokenizer.from_pretrained(llm_name)
20
- model = AutoModelForCausalLM.from_pretrained(llm_name)
21
 
 
 
22
 
23
- SYSTEM_PROMPT = "你是冯氏家谱机器人,只能根据给定资料回答,没有记载就说未记载。"
 
 
24
 
 
 
 
25
 
26
- def ask(question):
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
- prompt = f"{SYSTEM_PROMPT}\n\n【家谱资料】\n{context}\n\n【问题】{question}\n回答:"
 
 
 
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)