Llamaindex踩坑记录
·
今天在使用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=公司的休假政策是什么?
更多推荐
所有评论(0)