今天在使用llamaindex开发时发现检索部分的一个严重问题,检索时的输入跟向量化时的内容一样,但是检索出来的socre却差别很大。

根据这个问题我去看了开发部分的代码,发现我们插入的方式是使用llamaindex框架的textnode的方式来进行插入,textnode插入时是可以携带metadata的。

这个时候其实我就已经开始怀疑metadata的数据是不是也被向量化了,但是考虑到llamaindex框架应该不会犯这个错误,所以我去看了以下llamaindex的源码,在源码中我发现进行llamaindex向量化时也是使用node中的text和metadata的数据进行向量化的,具体如下:

textnode是继承于BaseNode,然后BaseNode下面有一个excluded_embed_metadata_keys参数用于排除metadata中无关的字段

具体实践请参考以下部分代码(重点在136行左右):

"""
QA搜索Score测试脚本

测试目的:验证使用 VectorStoreIndex.insert_nodes 插入数据后的搜索score

使用方法:
    python test_qa_score_normalize.py
"""

import sys
import os

# 添加项目根目录到路径
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from llama_index.core import Settings, StorageContext, VectorStoreIndex
from llama_index.core.schema import TextNode
from llama_index.embeddings.openai import OpenAIEmbedding
from llama_index.vector_stores.milvus import MilvusVectorStore
from pymilvus import connections, utility

# 导入项目配置
from app.config import MILVUS_HOST, MILVUS_PORT, VECTOR_DIMENSION, Env


def get_project_embedding_model() -> OpenAIEmbedding:
    """获取项目配置的Embedding模型"""
    embedding_model_full = Env.DEFAULT_EMBEDDING_MODEL

    if "." in embedding_model_full:
        provider, model_name = embedding_model_full.split(".", 1)
    else:
        provider = "openai"
        model_name = embedding_model_full

    # 提供商配置映射
    base_url_map = {
        "openai": Env.OPENAI_BASE_URL,
        "baai": Env.BAAI_BASE_URL,
        "qwen": Env.QWEN_BASE_URL,
        "azure": Env.AZURE_BASE_URL,
        "zhipu": Env.ZHIPU_BASE_URL,
    }
    api_key_map = {
        "openai": Env.OPENAI_API_KEY,
        "baai": Env.BAAI_API_KEY,
        "qwen": Env.QWEN_API_KEY,
        "azure": Env.AZURE_API_KEY,
        "zhipu": Env.ZHIPU_API_KEY,
    }

    return OpenAIEmbedding(
        model=model_name,
        api_key=api_key_map.get(provider, Env.OPENAI_API_KEY),
        api_base=base_url_map.get(provider, Env.OPENAI_BASE_URL),
        dimensions=VECTOR_DIMENSION,
    )


def cleanup_collection(collection_name: str):
    """清理测试collection"""
    try:
        connections.connect(host=MILVUS_HOST, port=MILVUS_PORT)
        if utility.has_collection(collection_name):
            utility.drop_collection(collection_name)
            print(f"  已清理collection: {collection_name}")
    except Exception as e:
        print(f"  清理collection失败: {e}")
    finally:
        try:
            connections.disconnect("default")
        except:
            pass


def run_test():
    """运行测试"""
    collection_name = "test_qa_score"

    # 测试数据
    test_texts = [
        "忘记密码怎么办?",
        "年假申请流程是什么?",
        "请假需要提前多久申请?",
        "公司的休假政策是什么?",
    ]

    # 查询文本(与第一条测试数据完全相同)
    query_text = "忘记密码怎么办?"

    print("\n" + "=" * 70)
    print("QA搜索Score测试")
    print("=" * 70)
    print(f"\n项目配置:")
    print(f"  Milvus地址: {MILVUS_HOST}:{MILVUS_PORT}")
    print(f"  向量维度: {VECTOR_DIMENSION}")
    print(f"  Embedding模型: {Env.DEFAULT_EMBEDDING_MODEL}")

    try:
        # 1. 初始化 embedding 模型
        print(f"\n初始化Embedding模型...")
        embed_model = get_project_embedding_model()
        Settings.embed_model = embed_model
        print(f"  模型初始化成功")

        # 2. 创建 Milvus 向量存储
        print(f"\n创建Milvus向量存储...")
        vector_store = MilvusVectorStore(
            uri=f"http://{MILVUS_HOST}:{MILVUS_PORT}",
            collection_name=collection_name,
            dim=VECTOR_DIMENSION,
            overwrite=True,
            similarity_metric="COSINE",
        )
        print(f"  向量存储创建成功")

        # 3. 创建索引
        print(f"\n创建VectorStoreIndex...")
        storage_context = StorageContext.from_defaults(vector_store=vector_store)
        index = VectorStoreIndex(
            nodes=[],
            storage_context=storage_context,
        )
        print(f"  索引创建成功")

        # 4. 创建节点并插入
        print(f"\n使用 insert_nodes 插入数据...")
        nodes = []
        for i, text in enumerate(test_texts):
            node_metadata = {
                "question_id": i,
                "org_id": i + 1
            }
            node = TextNode(
                text=text,
                metadata=node_metadata,
                # 如果对138-139进行注释就会导致向量匹配的结果出现小于0.99的情况
                # 添加排除在向量生成中的元数据字段
                excluded_embed_metadata_keys=["question_id", "org_id"],
                id_=f"test_{i}",
            )
            nodes.append(node)
            print(f"  创建节点: {text}")

        index.insert_nodes(nodes)
        print(f"  已插入 {len(nodes)} 个节点")

        # 5. 执行搜索
        print(f"\n执行搜索...")
        print(f"  查询文本: {query_text}")

        retriever = index.as_retriever(similarity_top_k=len(test_texts))
        search_results = retriever.retrieve(query_text)

        # 6. 输出结果
        print(f"\n搜索结果:")
        for i, hit in enumerate(search_results):
            print(f"  [{i + 1}] score={hit.score:.6f}, text={hit.text}")

        # 7. 分析
        print(f"\n" + "-" * 50)
        print("分析:")
        print("-" * 50)
        exact_match = next((r for r in search_results if r.text == query_text), None)
        if exact_match:
            print(f"  完全匹配文本的score: {exact_match.score:.6f}")
            if exact_match.score >= 0.99:
                print(f"  结论: score接近1.0,搜索正常")
            else:
                print(f"  结论: score偏低,可能存在问题")

    except Exception as e:
        print(f"\n测试失败: {e}")
        import traceback
        traceback.print_exc()

    finally:
        cleanup_collection(collection_name)


if __name__ == "__main__":
    print("=" * 70)
    print("开始测试...")
    print("=" * 70)

    run_test()

    print("\n" + "=" * 70)
    print("测试完成!")
    print("=" * 70)

输出结果:

# 注释前

使用 insert_nodes 插入数据...
  创建节点: 忘记密码怎么办?
  创建节点: 年假申请流程是什么?
  创建节点: 请假需要提前多久申请?
  创建节点: 公司的休假政策是什么?
  已插入 4 个节点

执行搜索...
  查询文本: 忘记密码怎么办?

搜索结果:
  [1] score=1.000000, text=忘记密码怎么办?
  [2] score=0.194886, text=请假需要提前多久申请?
  [3] score=0.154680, text=年假申请流程是什么?
  [4] score=0.121188, text=公司的休假政策是什么?

# 注释后:
使用 insert_nodes 插入数据...
  创建节点: 忘记密码怎么办?
  创建节点: 年假申请流程是什么?
  创建节点: 请假需要提前多久申请?
  创建节点: 公司的休假政策是什么?
  已插入 4 个节点

执行搜索...
  查询文本: 忘记密码怎么办?

搜索结果:
  [1] score=0.806405, text=忘记密码怎么办?
  [2] score=0.193372, text=年假申请流程是什么?
  [3] score=0.189418, text=请假需要提前多久申请?
  [4] score=0.146663, text=公司的休假政策是什么?
Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐