Chroma 教程

返回结果为什么是"列式"的

🎉摘要:本文详解Chroma数据库的get和query查询返回的列式(column-major)数据结构,提供使用zip遍历及自定义flatten_query_result函数将结果转为对象列表的实用方法,帮助开发者快速适应Chroma的查询结果格式。

作为一个经常使用 MySQL 数据库的你,首次使用 Chroma 的 get 和 query 查询数据,然而返回的不是“一个对象列表”,而是“一个字典,每个字段是一个列表”, 是不是有点懵。官方管这叫 column-major(列式),刚开始不习惯。

使用 get 查询数据

GetResult 的结构如下:

{
  "ids": ["id1", "id2", "id3"], # 所有的ID列,一个列表
  "embeddings": [[...], [...], [...]],   # 可选
  "documents": ["文本1", "文本2", "文本3"],
  "metadatas": [{...}, {...}, {...}],
  "uris": ["s3://a", "s3://b", "s3://c"],
  "data": None,
  "included": ["metadatas", "documents"],
}

同一条记录的信息,需要靠下标对齐。所以要遍历得用 zip( 用来把多个序列,按位置配对打包,拼成一组一组的元组),例如:

import chromadb

# 内存临时客户端
client = chromadb.EphemeralClient()
col = client.get_or_create_collection("demo_col")
# 添加几条数据,后面的查询/过滤才看得到东西
col.add(
    ids=["r1", "r2", "r3", "r4", "r5"],
    documents=[
        "退款流程:在订单页点申请退款,3 个工作日内到账",
        "发票开具:在订单详情页申请电子发票",
        "物流查询:发货后可在物流详情里看到进度",
        "退货政策:7 天无理由退货,运费买家承担",
        "账号安全:如何修改密码和绑定手机",
    ],
    metadatas=[
        {"cat": "售后", "year": 2024, "score": 0.9, "tags": ["退款"]},
        {"cat": "财务", "year": 2023, "score": 0.5, "tags": ["发票"]},
        {"cat": "物流", "year": 2024, "score": 0.7, "tags": ["物流"]},
        {"cat": "售后", "year": 2022, "score": 0.3, "tags": ["退货"]},
        {"cat": "账号", "year": 2024, "score": 0.6, "tags": ["安全"]},
    ],
)

# 拼接后,如:
# [(id,document,metadata),(...),...]
res = col.get(include=["documents", "metadatas"])
for id_, doc, meta in zip(res["ids"], res["documents"], res["metadatas"]):
    print(f"id={id_}, meta={meta}, doc={doc[:30]}")

运行上面代码,输出如下:

id=r1, meta={'year': 2024, 'score': 0.9, 'tags': ['退款'], 'cat': '售后'}, doc=退款流程:在订单页点申请退款,3 个工作日内到账
id=r2, meta={'year': 2023, 'cat': '财务', 'tags': ['发票'], 'score': 0.5}, doc=发票开具:在订单详情页申请电子发票
id=r3, meta={'tags': ['物流'], 'score': 0.7, 'year': 2024, 'cat': '物流'}, doc=物流查询:发货后可在物流详情里看到进度
id=r4, meta={'score': 0.3, 'cat': '售后', 'year': 2022, 'tags': ['退货']}, doc=退货政策:7 天无理由退货,运费买家承担
id=r5, meta={'year': 2024, 'score': 0.6, 'tags': ['安全'], 'cat': '账号'}, doc=账号安全:如何修改密码和绑定手机

使用 query 查询数据

QueryResult 要再套一层,因为 query 是批量的,你传了几个 query,就返回几组结果,结构如下:

{
  "ids": [
    ["id2", "id7"],  # 第 1 个 query 的结果
    ["id3", "id1"]   # 第 2 个 query 的结果
  ],
  "documents": [["文本2", "文本7"], ["文本3", "文本1"]],
  "metadatas": [[{...}, {...}], [{...}, {...}]],
  "distances": [[0.12, 0.34], [0.21, 0.55]],
  "included": ["metadatas", "documents", "distances"],
}

遍历要两层,示例:

import chromadb

client = chromadb.EphemeralClient()
col = client.get_or_create_collection("demo_col")
col.add(
    ids=["r1", "r2", "r3", "r4", "r5"],
    documents=[
        "退款流程:在订单页点申请退款,3 个工作日内到账",
        "发票开具:在订单详情页申请电子发票",
        "物流查询:发货后可在物流详情里看到进度",
        "退货政策:7 天无理由退货,运费买家承担",
        "账号安全:如何修改密码和绑定手机",
    ],
    metadatas=[
        {"cat": "售后", "year": 2024, "score": 0.9, "tags": ["退款"]},
        {"cat": "财务", "year": 2023, "score": 0.5, "tags": ["发票"]},
        {"cat": "物流", "year": 2024, "score": 0.7, "tags": ["物流"]},
        {"cat": "售后", "year": 2022, "score": 0.3, "tags": ["退货"]},
        {"cat": "账号", "year": 2024, "score": 0.6, "tags": ["安全"]},
    ],
)

res = col.query(query_texts=["问题一", "问题二"], n_results=2)
# 外层:每个 query 一组
for ids, docs, dists in zip(res["ids"], res["documents"], res["distances"]):
    # 内层:这一组里的每条
    for id_, doc, dist in zip(ids, docs, dists):
        print(id_, dist, doc)

运行示例,输出如下:

r2 1.1664421558380127 发票开具:在订单详情页申请电子发票
r5 1.199986219406128 账号安全:如何修改密码和绑定手机
r2 1.0322620868682861 发票开具:在订单详情页申请电子发票
r1 1.100636601448059 退款流程:在订单页点申请退款,3 个工作日内到账

如果你只想查一个问题,取第 0 组就行,就像这样 res["documents"][0]。

为了方便,我写了个小工具函数一直在用,省得每次都写两层循环:

def flatten_query_result(res):
    """把 query 的列式结果摊平成 [{'id':.., 'document':.., 'distance':.., 'metadata':..}]"""
    out = []
    for i, ids in enumerate(res["ids"]):
        for j, id_ in enumerate(ids):
            out.append({
                "id": id_,
                "document": res["documents"][i][j] if res.get("documents") else None,
                "metadata": res["metadatas"][i][j] if res.get("metadatas") else None,
                "distance": res["distances"][i][j] if res.get("distances") else None,
            })
    return out

用法如下:

res = col.query(query_texts=["问题一", "问题二"], n_results=2)
results = flatten_query_result(res)
for result in results:
    print(f"id={result["id"]}, document={result["document"]},"
          f" metadata={result["metadata"]}, distance={result["distance"]}")

输出如下:

id=r2, document=发票开具:在订单详情页申请电子发票, metadata={'tags': ['发票'], 'score': 0.5, 'year': 2023, 'cat': '财务'}, distance=1.1664421558380127
id=r5, document=账号安全:如何修改密码和绑定手机, metadata={'score': 0.6, 'year': 2024, 'cat': '账号', 'tags': ['安全']}, distance=1.199986219406128
id=r2, document=发票开具:在订单详情页申请电子发票, metadata={'score': 0.5, 'cat': '财务', 'year': 2023, 'tags': ['发票']}, distance=1.0322620868682861
id=r1, document=退款流程:在订单页点申请退款,3 个工作日内到账, metadata={'cat': '售后', 'year': 2024, 'tags': ['退款'], 'score': 0.9}, distance=1.100636601448059

建议在项目中创建工具库,将上面工具放入其中,用时导入调用即可。

说说我的看法
全部评论(
没有评论
关于
本网站专注于 Java、数据库(MySQL、Oracle)、Linux、软件架构及大数据等多领域技术知识分享。涵盖丰富的原创与精选技术文章,助力技术传播与交流。无论是技术新手渴望入门,还是资深开发者寻求进阶,这里都能为您提供深度见解与实用经验,让复杂编码变得轻松易懂,携手共赴技术提升新高度。如有侵权,请来信告知:hxstrive@outlook.com
其他应用
公众号