code/04-rag/vector_search.py

81 lignes · 3.4 Ko

Le 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]}")