code/04-rag/vector_search.py
81 lignes · 3.4 KoLe code et les sorties des programmes sont reproduits tels qu’ils ont tourné : commentaires et sorties sont donc en chinois.
"""从零写一个向量检索:把文档块变成向量存成矩阵,查询时算相似度取前 k 个。
再用 eval_qa.jsonl 里的 20 道题评估检索效果。
在 AI-Course 目录下运行:
python code/04-rag/vector_search.py 默认用 bge-small-zh-v1.5
python code/04-rag/vector_search.py intfloat/multilingual-e5-small 换一个嵌入模型
国内下载模型慢的话,先设置 export HF_ENDPOINT=https://hf-mirror.com
"""
import json
import sys
import time
from pathlib import Path
import numpy as np
from sentence_transformers import SentenceTransformer
sys.path.insert(0, str(Path(__file__).parent))
from chunking import load_docs, split_by_heading_capped
HERE = Path(__file__).parent
class VectorIndex:
def __init__(self, model_name):
self.model = SentenceTransformer(model_name)
# e5 系列要求给查询和文档分别加上前缀,这是它训练时的约定
self.q_prefix, self.d_prefix = ("query: ", "passage: ") if "e5" in model_name else ("", "")
self.chunks = [] # [(文件名, 文本), ...]
self.matrix = None # 每一行是一个块的向量
def build(self, chunks):
self.chunks = chunks
texts = [self.d_prefix + text for _, text in chunks]
self.matrix = self.model.encode(texts, normalize_embeddings=True, batch_size=32)
def search(self, query, k=5):
q = self.model.encode([self.q_prefix + query], normalize_embeddings=True)[0]
scores = self.matrix @ q # 向量都归一化过了,点积就是余弦相似度
top = np.argsort(-scores)[:k]
return [(float(scores[i]), *self.chunks[i]) for i in top]
def is_hit(result, qa):
_, file, text = result
return file == qa["file"] and qa["keyword"] in text
def evaluate(search, questions, k=5):
"""返回第 1 名命中率、前 3 名命中率、前 5 名命中率、MRR,以及没找到的题。"""
ranks, misses = [], []
for qa in questions:
results = search(qa["question"], k)
rank = next((i + 1 for i, r in enumerate(results) if is_hit(r, qa)), None)
ranks.append(rank)
if rank is None:
misses.append((qa, results[0]))
n = len(questions)
hit = lambda top: sum(1 for r in ranks if r and r <= top) / n
mrr = sum(1 / r for r in ranks if r) / n
return hit(1), hit(3), hit(5), mrr, misses
if __name__ == "__main__":
MODEL_NAME = sys.argv[1] if len(sys.argv) > 1 else "BAAI/bge-small-zh-v1.5"
chunks =[(file, c) for file, text in load_docs().items() for c in split_by_heading_capped(text)]
questions = [json.loads(line) for line in (HERE / "eval_qa.jsonl").read_text().splitlines()]
start = time.time()
index = VectorIndex(MODEL_NAME)
index.build(chunks)
print(f"{MODEL_NAME}:{len(chunks)} 个块,向量 {index.matrix.shape},建索引用了 {time.time() - start:.1f} 秒")
score, file, text = index.search("怎么关闭 SSL 证书校验?", k=1)[0]
print(f"\n示例:「怎么关闭 SSL 证书校验?」最相似的块来自 {file},相似度 {score:.3f}")
print(text[:200])
h1, h3, h5, mrr, misses = evaluate(index.search, questions)
print(f"\n20 道题:第 1 名命中 {h1:.0%},前 3 名命中 {h3:.0%},前 5 名命中 {h5:.0%},MRR {mrr:.3f}")
for qa, (score, file, text) in misses:
print(f" 没找到:{qa['question']}(应在 {qa['file']})→ 第 1 名是 {file}:{text.splitlines()[0][:40]}")