Embedding原理简介

Embedding(嵌入)是将文本转换为高维向量的过程。语义相近的文本,其向量在高维空间中的距离也更近。RAG系统依赖Embedding的质量来判断"哪些文档片段与用户问题最相关"。

Embedding模型的训练目标是:让语义相似的句子对(正样本)的向量余弦相似度高,让不相关句子对(负样本)的相似度低。常见的训练方式是对比学习(Contrastive Learning)。

主流Embedding模型对比

模型 维度 中文支持 调用方式 成本 适用场景
OpenAI text-embedding-3-small 1536(可压缩) 良好 API $0.02/1M tokens 通用,快速上手
OpenAI text-embedding-3-large 3072(可压缩) 良好 API $0.13/1M tokens 高精度需求
BGE-large-zh-v1.5 1024 优秀 本地部署 免费 中文文档,私有部署
M3E-large 1024 优秀 本地部署 免费 中文,轻量级场景
Jina-embeddings-v3 1024 良好 API/本地 有免费额度 多语言,长文档
BGE-M3 1024 优秀 本地部署 免费 多语言,多粒度检索

中文文档该用哪个?

对于以中文为主的文档库,推荐优先顺序如下:

调用代码示例

BGE本地部署(使用sentence-transformers)

from sentence_transformers import SentenceTransformer

# 自动下载模型(约1.3GB)
model = SentenceTransformer("BAAI/bge-large-zh-v1.5")

texts = [
    "RAG技术可以有效减少大模型的幻觉",
    "向量数据库是RAG系统的核心组件",
]

# 对中文查询,BGE需要添加指令前缀
instruction = "为这个句子生成表示以用于检索相关文章:"
query = instruction + "如何减少AI幻觉?"

query_embedding = model.encode(query, normalize_embeddings=True)
doc_embeddings = model.encode(texts, normalize_embeddings=True)

# 计算余弦相似度
import numpy as np
similarities = np.dot(doc_embeddings, query_embedding)
print(similarities)  # [0.85, 0.72]

OpenAI Embedding API

from openai import OpenAI

client = OpenAI()

def embed(texts: list[str]) -> list[list[float]]:
    resp = client.embeddings.create(
        model="text-embedding-3-small",
        input=texts,
        # 可以将维度压缩到256,减少存储
        dimensions=512
    )
    return [item.embedding for item in resp.data]

embeddings = embed(["RAG原理", "向量检索"])
print(f"维度:{len(embeddings[0])}")  # 512

维度选择与存储成本

OpenAI text-embedding-3系列支持维度压缩(Matryoshka Representation Learning),可以在调用时指定较小的维度(如256、512),在牺牲少量精度的情况下大幅降低存储和检索成本。

对于大多数场景,512维已经足够,无需使用全维度。

Embedding质量评估

评估方法:构建一个问答对测试集(100-200条),统计正确答案在检索Top-K中的命中率。

def evaluate_retrieval(qa_pairs, retriever, k=5):
    """
    qa_pairs: [(question, expected_doc_id), ...]
    """
    hits = 0
    for question, expected_id in qa_pairs:
        results = retriever.search(question, k=k)
        retrieved_ids = [r.id for r in results]
        if expected_id in retrieved_ids:
            hits += 1
    return hits / len(qa_pairs)  # Recall@K

recall_at_5 = evaluate_retrieval(test_set, my_retriever, k=5)
print(f"Recall@5: {recall_at_5:.2%}")

通过在测试集上对比不同Embedding模型的Recall@5,可以客观选出最适合你数据的模型。