You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Tensorflow Recommender:向ScaNN传入查询嵌入的数据类型咨询

ScaNN直接传入查询嵌入的数据类型说明

ScaNN完全支持直接传入查询嵌入而非模型,核心要求是float32类型的二维numpy数组,具体细节如下:

  • 数据类型强制要求:所有嵌入(查询、候选)必须是np.float32类型,其他浮点类型(如float64)可能导致不兼容。
  • 形状要求:
    • 单条查询嵌入:必须转为二维数组,比如你的示例[1, 0.3, 0.4]要处理成np.array([[1, 0.3, 0.4]], dtype=np.float32)
    • 候选嵌入:保持(候选数量, 嵌入维度)的二维数组结构,如你的示例转成np.array([[0.2, 1, 0.4],[0.3,0.1,0.56]], dtype=np.float32)

完整代码示例

import numpy as np
import scann

# 预处理候选嵌入
candidate_embeddings = np.array([[0.2, 1, 0.4], [0.3, 0.1, 0.56]], dtype=np.float32)

# 构建ScaNN索引(以点积相似度为例)
searcher = scann.scann_ops_pybind.builder(
    candidate_embeddings, 
    num_neighbors=2, 
    distance_measure="dot_product"
).tree(
    num_leaves=2, 
    num_leaves_to_search=2, 
    training_sample_size=2
).score_ah(
    2, 
    anisotropic_quantization_threshold=0.2
).build()

# 预处理查询嵌入(转为二维float32数组)
query_embedding = np.array([[1, 0.3, 0.4]], dtype=np.float32)

# 执行搜索
neighbor_indices, similarity_scores = searcher.search(query_embedding)
print("匹配候选索引:", neighbor_indices)
print("相似度分数:", similarity_scores)

常见问题排查

  • 如果之前传入一维numpy数组报错:ScaNN的search方法默认期望批量查询输入,必须是二维数组,哪怕只有一条查询。
  • 如果数据类型不匹配报错:直接用dtype=np.float32显式指定类型,避免默认的float64。

内容的提问来源于stack exchange,提问作者uhXAj1

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 05:10:28