projects/repobot/v4/retrieval.py

165 行 · 6.3 KB

コードと実行結果は実際に動かしたときのまま載せているため、コメントと出力は中国語です。

"""检索的部分:切分、BM25、向量检索、RRF 融合、可选的重排,以及查询改写。
每一块在课程 04 模块里都单独讲过,这里只是把它们放在一起。
"""
import hashlib
import json
import math
import re
import threading
from collections import Counter
from pathlib import Path

import numpy as np

import llm

EMBED_MODEL = "intfloat/multilingual-e5-small"
RERANK_MODEL = "BAAI/bge-reranker-base"
_MODEL_LOCK = threading.Lock()  # 同一时间只让一个线程调用本地的嵌入模型和重排模型


# ---------- 切分(04 模块第 2 课) ----------

def load_docs(docs_dir):
    docs_dir = Path(docs_dir)
    return {str(p.relative_to(docs_dir)): p.read_text() for p in sorted(docs_dir.rglob("*.md")) if p.name != "LICENSE.md"}


def split_by_heading(text):
    chunks, current, in_code = [], [], False
    for line in text.splitlines():
        if line.startswith("```"):
            in_code = not in_code
        if not in_code and re.match(r"#{1,3} ", line) and current:
            chunks.append("\n".join(current).strip())
            current = []
        current.append(line)
    if current:
        chunks.append("\n".join(current).strip())
    return [c for c in chunks if c]


def split_capped(text, max_size=1500):
    chunks = []
    for section in split_by_heading(text):
        if len(section) <= max_size:
            chunks.append(section)
            continue
        title = section.splitlines()[0] if section.startswith("#") else ""
        current = ""
        for para in section.split("\n\n"):
            if current and len(current) + len(para) > max_size:
                chunks.append(current.strip())
                current = title + "\n\n" if title else ""
            current += para + "\n\n"
        if current.strip():
            chunks.append(current.strip())
    return chunks


# ---------- BM25 和 RRF(04 模块第 4 课) ----------

def tokenize(text):
    return re.findall(r"[a-z0-9_]+|[一-鿿]", text.lower())


class BM25:
    def __init__(self, texts, k1=1.5, b=0.75):
        self.docs = [tokenize(t) for t in texts]
        self.avg_len = sum(len(d) for d in self.docs) / len(self.docs)
        self.tf = [Counter(d) for d in self.docs]
        df = Counter(w for d in self.docs for w in set(d))
        n = len(self.docs)
        self.idf = {w: math.log(1 + (n - c + 0.5) / (c + 0.5)) for w, c in df.items()}
        self.k1, self.b = k1, b

    def search(self, query, k):
        words = tokenize(query)
        scores = []
        for i, (tf, doc) in enumerate(zip(self.tf, self.docs)):
            s = sum(self.idf[w] * tf[w] * (self.k1 + 1) /
                    (tf[w] + self.k1 * (1 - self.b + self.b * len(doc) / self.avg_len)) for w in words if w in tf)
            scores.append((s, i))
        scores.sort(reverse=True)
        return [i for _, i in scores[:k]]


def rrf(rank_lists, c=60):
    scores = Counter()
    for ranks in rank_lists:
        for r, i in enumerate(ranks, 1):
            scores[i] += 1 / (c + r)
    return [i for i, _ in scores.most_common()]


# ---------- 检索器 ----------

class Retriever:
    def __init__(self, docs_dir, cache_dir, rerank=False):
        from sentence_transformers import SentenceTransformer

        self.chunks = [(f, c) for f, t in load_docs(docs_dir).items() for c in split_capped(t)]
        texts = [c for _, c in self.chunks]
        self.bm25 = BM25(texts)
        self.embedder = SentenceTransformer(EMBED_MODEL)

        # 文档或模型不变,就直接读缓存的向量,不用每次启动都重新计算
        cache_dir = Path(cache_dir)
        cache_dir.mkdir(exist_ok=True)
        key = hashlib.sha256((EMBED_MODEL + json.dumps(self.chunks)).encode()).hexdigest()[:16]
        path = cache_dir / f"vectors-{key}.npy"
        if path.exists():
            self.matrix = np.load(path)
        else:
            self.matrix = self.embedder.encode(["passage: " + t for t in texts], normalize_embeddings=True, batch_size=32)
            np.save(path, self.matrix)

        self.reranker = None
        if rerank:
            from sentence_transformers import CrossEncoder
            self.reranker = CrossEncoder(RERANK_MODEL, max_length=512)

    def vector_search(self, query, k):
        # 多个线程同时调用本地模型时,PyTorch 各自再开一堆线程,会互相争抢到几乎卡死。
        # 一次只让一个线程算向量,算一个问题只要几毫秒,排队几乎不影响速度
        with _MODEL_LOCK:
            q = self.embedder.encode(["query: " + query], normalize_embeddings=True)[0]
        return list(np.argsort(-(self.matrix @ q))[:k])

    def search(self, query, k=5):
        """返回 [(文件名, 文本), ...]。先两路各取 20 个用 RRF 融合,有重排模型就再精排一次。"""
        pool = rrf([self.vector_search(query, 20), self.bm25.search(query, 20)])[:20]
        if self.reranker is not None:
            with _MODEL_LOCK:
                scores = self.reranker.predict([(query, self.chunks[i][1]) for i in pool])
            pool = [pool[j] for j in np.argsort(-scores)]
        return [self.chunks[i] for i in pool[:k]]


# ---------- 查询改写 ----------

REWRITE_PROMPT = """你要为一个 httpx 答疑助手生成文档检索词。httpx 的文档是英文的。

根据"最近的对话"理解用户"最新的问题"到底在问什么(比如"那异步呢"要结合上文补全),
然后输出一行英文检索词:包含问题的完整英文表述,以及文档里可能出现的参数名、类名、术语。只输出这一行。"""


class QueryRewriter:
    def __init__(self, cache_dir):
        self.path = Path(cache_dir) / "rewrites.json"
        self.cache = json.loads(self.path.read_text()) if self.path.exists() else {}
        self.last_usage = None

    def rewrite(self, question, history=()):
        recent = "\n".join(f"{m['role']}: {m['content'][:300]}" for m in list(history)[-4:])
        key = recent + "\n>>> " + question
        self.last_usage = None
        if key not in self.cache:
            text, self.last_usage = llm.chat([
                {"role": "system", "content": REWRITE_PROMPT},
                {"role": "user", "content": f"最近的对话:\n{recent or '(无)'}\n\n最新的问题:{question}"},
            ])
            self.cache[key] = text.strip()
            self.path.write_text(json.dumps(self.cache, ensure_ascii=False, indent=2))
        return self.cache[key]