"""不依赖模型服务的检索与融合基线。"""
from collections import defaultdict
import numpy as np


def cosine_topk(query, docs, k=5):
    q = query / (np.linalg.norm(query) + 1e-12)
    d = docs / (np.linalg.norm(docs, axis=1, keepdims=True) + 1e-12)
    scores = d @ q; k = min(k, len(scores))
    ids = np.argpartition(-scores, k-1)[:k]
    ids = ids[np.argsort(-scores[ids])]
    return [(int(i), float(scores[i])) for i in ids]


def rrf(rank_lists, constant=60):
    score = defaultdict(float)
    for results in rank_lists:
        for rank, doc_id in enumerate(results, 1):
            score[doc_id] += 1 / (constant + rank)
    return sorted(score.items(), key=lambda x: x[1], reverse=True)


if __name__ == "__main__":
    rng = np.random.default_rng(0)
    hits = cosine_topk(rng.normal(size=8), rng.normal(size=(20, 8)), 3)
    print("dense:", hits)
    print("rrf:", rrf([[1, 2, 3], [2, 4, 1]]))

