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
相关产品推荐
相关产品推荐

