RAG 基础

2026-04-21

Retrievel-Augmented Genration 检索 增强 生成

先从资料库检索相关内容,再基于内容生成答案

常用于智能客服或知识库

参考内容open in new window

背景

如果将完整资料交给模型

  • 由于模型拥有有限的上下文,如果提供资料超过上下文大小,会出现读了后面,忘了前面。
  • 且由于阅读量较大,每次回答都需阅读大量资料,推理成本过高
  • 推理速度也会很慢

将资料分片,每次从分片中获取与问题相关的片段,将这些片段与问题发送给大模型,实现更高效的问题解答

基本流程

RAG 的基本流程分为两个阶段。对知识库文档先进行解析存储,即预处理阶段。后续在线问答阶段,则是根据用户的问题,在已存储的文档中查询语义相近的文档结果,根据内容做出回答

预处理

  • 分片: 拥有多种方式,如 按指定字数切分、按段落、按章节、按页码
  • 索引: 共两个步骤
    • 通过 Embedding 将片段文本转换成向量。语句通过 Embedding 模型open in new window 转换成向量,含义相近的语句会解析成相近的向量
    • 将片段文本与向量存储进向量数据库中。向量数据库可用于存储与查询向量。后续得到向量后,查询相近的向量与对应的原始文本,从而实现后续发送给大模型。通常有两列,一列是原始文本,一列是解析得到的向量

回答

  • 召回: 搜索与用户问题相关的片段
    • 将用户问题发送给 Embedding 模型,得到对应的向量
    • 通过向量查询向量数据库,得到固定个数的最相似结果(计算向量相似度,成本低,准确率低,做初筛)
      • 余弦相似度: 计算相关夹角的 cos 值,即计算夹角大小,不关注向量长度的不同(仅关注方向,忽略强度)。该算法避免数量级带来的差异,
      • 欧式距离: 计算向量的距离
      • 点积: 计算向量 A 在向量 B 投影长度与向量 B 长度的乘积,越大相似度越高
  • 重排: 重新排序。从召回得到的结果中,进行再次筛选,挑选出指定个数最相似的结果(使用 cross-encoder 模型计算与问题的相似度,成本高,准确率高,做复筛)
  • 生成: 将重排结果与问题发送给大模型,大模型根据提供内容生成结果返回

代码实现

安装依赖

使用 python 的 uv 进行包管理,安装相应的依赖

uv add sentence_transformers chromadb google-genai python-dotenv

以下是各个包的简介

  • sentence_transformers: embedding 模型 与 cross-encoder 模型
  • chromadb: 向量数据库
  • google-genai: 调用 google gemini 的依赖,此处使用 gemini 作为生成阶段的大模型,因此需要安装此包
  • python-dotenv: 环境变量管理

分片

分片有多种分割方式,如按行、按字数、按章节等。这里介绍两种分片方式的实现

  • 按行切分

    def split_into_chunks(doc_file: str) -> list[str]:
        with open(doc_file, 'r') as file:
            content = file.read()
        # 按行进行切分
        return [chunk for chunk in content.split('\n\n')]
    
  • 按字数切分。保证每个块有前后块的相邻字符,减少切分产生丢失语义的可能(经专业机构测试,10%-20%效果最优)

    # 传入原文,切分窗口大小(400),保留相邻块的字符数(100)
    def split_text(text: str, chunk_size: int = CHUNK_SIZE, overlap: int = CHUNK_OVERLAP):
        text = text.strip()
        if not text:
            return []
        chunks = []
        start = 0
        while start < len(text):
            end = start + chunk_size
            chunks.append(text[start:end])
            start = end - overlap  # 跳转开始位置,重复一定量的字符
        return chunk
    

索引

计算每个文本段的向量,将其进行存储

  • 文本转换成向量
    • 使用 sentence_transformers 解析向量。预训练好的通用模型,向量维度和语义空间是固定不变的,计算单条文本完全独立

      from sentence_transformers import SentenceTransformer
      
      # 加载一个 embedding 模型 - 会从网上拉取大模型,要保证网络通畅
      embedding_model = SentenceTransformer('shibing624/text2vec-base-chinese')
      
      # 将文本转换成向量
      def embed_chunk(chunk: str) -> list[float]:
          embedding = embedding_model.encode(chunk)
          return embedding.tolist()
      
    • 使用 scikit-learn 解析向量。TfidfVectorizer(传统统计学/稀疏向量):基于当前数据集的统计规律(词频 + 逆文档频率),每个维度的含义和权重取决于整个语料库。

      import pickle
      from sklearn.feature_extraction.text import TfidfVectorizer
      
      VECTORIZER_PATH = os.path.join(DATA_DIR, 'vectorizer.pkl')  # 存储向量解析
      # 向量解析器
      _vectorizer = None
      
      def get_vectorizer() -> TfidfVectorizer:
          global _vectorizer
          if not _vectorizer:
              # 磁盘中已经存在
              if os.path.exists(VECTORIZER_PATH):
                  with open(VECTORIZER_PATH, 'rb') as f:
                      _vectorizer = pickle.load(f)
              else:
                  _vectorizer = TfidfVectorizer(
                      analyzer='char_wb', 
                      ngram_range=(2, 4),
                      max_features=10000
                  )
          return _vectorizer
      
      # 传入多个块,得到多个块的向量
      def add_new_chunks(chunks: list[str]) -> list[list[f]]:
          vec = get_vectorizer()
          # 获取文本的词汇表
          vec.fit(chunks)
          # 解析向量
          embeddings = vec.transform(chunks)
      
          ret = [embedding.toarray()[0].tolist() for embedding in embeddings]
      
          # IDF 的计算必须依赖“全局语料”(文档全局统计). TF-IDF=TF(词频)×IDF(逆文档频率)
          # 刷新并加载向量器。每次添加新的文本,需要与旧文本一起训练。保证新文本与旧文本中,相近的语义,生成的向量也相近
          # 用全量语料重建向量器,保证所有文档的向量在同一词汇空间
          rebuild_vectorizer()
          _re_encode_all()
      
      # 在全量 chunk 上重新训练向量器
      def rebuild_vectorizer():
          global _vectorizer
          all_chunks = storage.get_all_chunks()
          # 从 chunks 的结构字典列表中获取 切片文本列表
          texts = [c['text'] for c in all_chunks]
          _vectorizer = TfidfVectorzier(
              analyzer='char_wb', 
              ngram_range=(2, 4),
              max_features=10000
          )
          _vectorizer.fit(texts)
          with open(VECTORIZER_PATH, 'wb') as f:
              pickle.dump(_vectorizer, f)
          return _vectorizer
      
      # 用新的全局向量器冲洗编码所有 chunk / 更新所有 chunk 的向量信息
      def _re_encode_all():
          all_chunks = storage.get_all_chunks()
          if not all_chunks:
              return
          vec = get_vectorizer()
          # 从 chunks 的结构字典列表中获取 切片文本列表
          texts = [c['text'] for c in all_chunks]
          matrices = vec.transform(texts)
          ds = storage.get_docstore()
          for i, row in enumerate(matrices):
              # 更新向量信息
              ds['chunks'][i]['embedding'] = row.toarray()[0].tolist()
          storage.save_docstore(ds)
      
  • 向量存储
    import chromadb
    # 数据库连接对象
    chromadb_client = chromadb.EphemeralClient()  # EphemeralClient 是内存型向量数据库,用于测试
    # chromadb_client = chromadb.PersistentClient('<文件路径>')  # 写入磁盘的数据库
    # 创建集合
    chromadb_collection = chromadb_client.get_or_create_collection(name='default')
    
    # 传入 原文数组 与 对应的多个向量 进行保存
    def save_embeddings(chunks: list[str], embeddings: list[list[float]]) -> None:
        # chromadb 要求每条数据需要有 ID,这里生成一个 ID 列表
        ids = [str(i) for i in range(len(chunks))]  
        # 存储入库
        chromadb_collection.add(
            documents=chunks,
            embeddings=embeddings,
            ids=ids
        )
    

Tip: 在现代检索(RAG、语义搜索)中,更推崇预训练的 Embedding 模型,正是因为它们在增量插入数据时只需要对新切片计算一次,旧切片的向量持久化后永久有效,工程维护成本极低。TF-IDF 虽然单次计算极快,但语料库越大,增量更新的代价(全量重算)越可怕

召回

将问题转换成向量,并查询语义相近的向量结果

  • 使用 embedding 模型进行召回

    # 传入问题与召回数量,得到多个召回结果
    def retrieve(query: str, top_k: int) -> list[str]:
        # 将问题解析成向量
        query_embedding = embed_chunk(query)
        # 根据向量在数据库中查询出结果
        results = chromadb_collection.query(
            query_embeddings=[qyery_embedding],
            n_results=top_k
        )
        return results['documents'][0]
    
  • 使用 scikit-learn 的 TfidfVectorizer 召回。将问题转换成 TF-IDF 向量,与所有 chunk 计算余弦相似度

    from sklearn.metrics.pairwise import cosine_similarity
    
    def retrieve(query: str, top_k: int = 4) -> list[dict]
        all_chunks = storage.get_all_chunks()
        if not all_chunks:
            return []
        vec = get_vectorizer()
        query_embedding = vec.tranform([query]).toarray()
        chunks_embeddings = np.array([c['embedding'] for c in all_chunks])
        # 计算余弦相似度
        sims = consine_similarity(query_embedding, chunks_embeddings)[0]
        ranked = sorted(zip(all_chunks, sims), key=lambda x: x[1], reverse=True)
        top = ranked[:top_k]
    
        return [chunk for chunk, _ in top]
    

重排

在使用召回得到语义相近的文本后,使用成本更高的 cross-encoder 模型再次计算这些文本与问题的相似度,重新进行排序。以提高回答的准确率

from sentence_transformers import CrossEncoder

# 传入 问题,召回结果,复筛个数,得到 重排结果列表
def rerank(query: str, retrieved_chunks: list[str], top_k: int) -> list[str]:
    # 创建 cross_encoder 模型 - 也会从网上拉取大模型,要保证网络通畅
    cross_encoder = CrossEncoder('cross-encoder/mmarco-mMiniLMv2-L12-H384-v1')
    # 创建问题与原文一一对应的列表
    pairs = [(query, chunk) for chunk in retrieved_chunks]
    # 计算每个召回的原文与问题的相似度
    scores = cross_encoder.predict(pairs)
    # 创建原文与评分一一对应的列表 [(原文, 分数), (原文, 分数)]
    chunk_with_score_list = [(chunk, score) for chunk, score in zip(retrieved_chunks, scores)]
    # 按照分数重新排序
    chunk_with_score_list.sort(key=lambda pair: pair[1], reverse=True)
    # 返回截取后的结果
    return [chunk for chunk, _ in chunk_with_score_list][:top_k]

生成

此处需要调用大模型,这里以 gemini 为例,在 Google AI Studioopen in new window 页面创建一个 API KEY

  • 创建 .env 文件并写入 API KEY
    GEMINI_API_KEY=*****
    
  • 在配置好 API KEY 后编写调用大模型的代码。构造提示词传递给大模型,让其生成回复结果
    from dotenv import load_dotenv
    from google import genai
    
    load_dotenv()  # .env 中的内容设置为环境变量
    
    # 创建 gemini 客户端,会从环境变量中读取 GEMINI_API_KEY 的值
    google_client = genai.Client()
    
    def generate(query: str, chunks: list[str]) -> str:
        # 构建输入大模型的 prompt
        prompt = f"""你是一位知识助手,请根据用户问题和下列片段生成准确的答案。
    
    用户问题: {query}
    
    相关片段: 
    {"\n\n".join(chunks)}
    
    请基于上述内容作答,不要编造信息。"""
        
        # 获取大模型根据输入的提示词得到的结果
        response = google_client.models.generate_content(
            model='gemini-2.5-flash',
            contents=prompt
        )
        return response.text
    

如果需要 SSE 流式返回,并让大模型记住历史消息,可参考以下代码

# 构造系统提示词 - 传入召回的文本内容
def _build_system_prompt(context_texts: list[str]) -> str:
    if not context_texts:
        return '你是一个知识库助手。当前知识库为空,请告知用户先上传文档'
    # 将所有召回的文本内容拼接到提示词中
    ctx = '\n---\n'.join(context_texts)
    return (
        '你是一个知识库助手。请仅根据以下资料回答用户的问题,'
        '如果资料中没有答案,请如实说”知识库中暂无相关内容“\n\n'
        f'【参考资料】\n{ctx}'
    )

# RAG 问题处理的 SSE 流式会话函数
def stream_chat(history_messages: list[str], query: str, top_k: int = RETRIEVAL_TOP_K) -> Iterator[str]:
    # 获取召回的 chunk 列表
    retrived = rag,retrieve(query, top_k)
    # 获取召回的原文内容
    context_texts = [c['text'] for c in retrieved]
    # 获取系统提示词
    system = _build_system_prompt(context_texts)
    
    # 组装历史消息
    messages = [{'role': 'system', 'context': system}]
    for m in history_messages:
        messages.append({'role': m['role'], 'content': m['content']})
    messages.append({'role': 'user', 'content': query})

    try:
        for token in chat_stream(messages):  # 此处调用发送大模型得到 stream 返回的 chat_stream 函数
            # 将返回的字符转换成 json 字符串,便于后续在前方拼接 "data:"
            payload = json.dumps({'type': 'content', 'text': token}, ensure_ascii=False)
            yield f'data: {payload}\n\n'
        yield f'data: [DONE]\n\n'
    except Exception as e:
        payload = json.dumps({'type': 'error', 'message': str(e), ensure_ascii=False})
        yield f'data: {payload}\n\n'

大模型对话后端接口

由于最终效果是用户通过对话框实现文本内容的检索,需要实现大模型对话功能

在接入大模型时,每个大模型接入的 SDK 不同,最终可使用 langchain 封装,实现修改极少量代码可接入不同的大模型

此处不对 langchain 做赘述。按照 deepseek 官方文档,介绍如何介入 deepseek。由于文档中说明可使用 openai 的 sdk 接入,安装 openai 的依赖

pip install openai

以下是 Flask 后端接口的实现流程

  • 对话模块:deepseek 兼容 openai 的规范,可使用 OpenAI 的 sdk 进行接入

    from openai import OpenAI
    
    _client = None
    
    def get_client() -> OpenAI:
        global _client
        if not _client:
            _client = OpenAI(
                api_key=API_KEY,
                base_url=BASE_URL
            )
        return _client
    
    # 传入记录会话消息历史的数组,每条记录包含发送人和内容的字典
    def chat_stream(messages: list[dict]) -> Iterator[str]
        stream = get_client().chat.completions.create(
            model=DEEPSEEK_MODEL,  # deepseek-v4-flash
            messages=messages,
            stream=True,  # 流式输出
            temperature=0.7,
            max_tokens=2048
        )
        for chunk in stream:
            delta = chunk.choices[0].delta  # 根据 ds 官网的接收格式,获取大模型返回内容
            if delta.content:  # 判断返回是否包含内容
                yield delta.content
    
  • 利用对话模块编写一个终端会话功能函数

    def chat_in_cli():
        print('deepseek 控制台对话')
        # 记录会话消息的数组
        messages = [
            {'role': 'system', 'content': '你是一个助手,请用中文回答问题'}  # 默认消息中第一条为系统提示词
        ]
        while True:
            user_input = input('\n你:').strip()  # 获取用户输入
            if not user_input:
                continue
            if user_input.lower() in ('exit', 'quit'):
                print('再见')
                break
            # 在记录消息记录的数组内追加用户输入
            messages.append({'role': 'user', 'content': user_input})
            print('助手: ',end='', flush=True)
            # 存储大模型返回的答案流
            full_anwser = ''
            for token in chat_stream(messages):  # 每次与大模型对话都传入完整的消息记录数组
                print(token, end='', flush=True)  # print 默认每次输出会换行,将换行取消并立即刷新
                full_anwser += token
            # 换行处理
            print()
            # 将大模型返回的答案放进消息记录数组内
            messages.append({'role': 'assistant', 'content': full_answer})
    
  • SSE 流式会话函数。大模型返回时,会一个字一个字返回。最终返回 [DONE] 表示大模型的此次内容输出已经结束

    # 固定格式 `data: <内容>`
    # data: {"type": "content","text": "你"}
    # data: {"type": "content","text": "好"}
    # data: [DONE]
    
    # 传入完整的消息记录数组与本次的问题
    def stream_chat(history_messages: list[dict], query: str) -> Iterator[str]:
        # 参考终端会话功能函数,用一个数组记录完整的消息历史
        messages = [
            {'role': 'system', 'content': '你是一个助手,请用中文回答问题'}  # 默认消息中第一条为系统提示词]
        ]
        for m in history_messages:
            messages.append({
                'role': m['role'], 'content': m['content']
            })
        # 放入本次问题
        messages.append({'role': 'user', 'content': query}) 
    
        try:
            for token in chat_stream(messages):
                # 将返回的字符转换成 json 字符串,便于后续在前方拼接 "data:"
                payload = json.dumps({'type': 'content', 'text': token}, ensure_ascii=False)
                yield f'data: {payload}\n\n'
            yield f'data: [DONE]\n\n'
        except Exception as e:
            payload = json.dumps({'type': 'error', 'message': str(e), ensure_ascii=False})
            yield f'data: {payload}\n\n'
    
  • 使用 SSE 流式会话的 Flask API 接口

    @app.route('/api/chat/stream', methods=['POST'])
    def chat_stream():
        data = request.get_json(force=True)
        query = data.get('message', '').strip()
        session_id = data.get('session_id', '').strip()  # 新会话为空,后端自行生成 id,防止后端查询会话时异常
        if not query:
            return jsonify({'error': 'empty message'}), 4000
        
        # 没 session id,说明是新会话,后端执行创建会话逻辑
        if not session_id or not storage.get_session(session_id):
            s = storage.create_session()
            session_id = s['id']
        
        # 定义一个内部方法,调用大模型,存储会话与聊天记录,返回大模型响应内容
        def generate() -> Iterator[str]:
            session = storage.get_session(session_id)
            # 获取该会话的完整消息历史记录
            history = session.get('message', []) if session esle []
            # 存储大模型的完整响应文本,用于存储到历史记录中
            full_answer = ''
            # 调用 SSE 流式会话函数
            for sse_line in stream_chat(history, query):
                # 存储大模型的完整响应内容
                if sse_line.strip():
                    # 某个字节出现问题时,跳过该字节
                    try:
                        raw = sse_line.replace('data: ', '').strip()
                        if raw and raw != '[DONE]':
                            d = json.loads(raw)
                            if d.get('type') == 'content':
                                full_anwser += d.get('text': '')
                    except Exception:
                        pass
                # 该函数直接返回 SSE 流
                yield sse_line
        
            if session:
                # 追加历史记录
                session['message'].append({'role': 'user', 'content': query})
                session['message'].append({'role': 'assistant', 'content': full_answer})
    
                # 没有标题或标题是系统自动自动生成的,创建一个新标题
                if not session.get('title') or session['title'].startswith('会话 '):
                    # 截取问题的前 30 字符作为新标题
                    session['title'] = query[:30]  + ('...' if len(query) > 30 else '')
                storage.save_session(session)
    
        return Response(
            generate(),  # 将 yield 的 SSE 流直接返回
            mimetype='text/event-stream',  # 设置相应内容为流式
            headers={
                'Cache-Control': 'no-cache',
                'Connection': 'keep-alive',
                'X-Accel-Buffering': 'no'
            }
        )