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

能否将候选数据集作为TensorFlow检索TopK模型的输入?

动态传入候选集到TensorFlow Retrieval模型的解决方案

BruteForce层不支持动态候选集

tfrs.layers.factorized_top_k.BruteForce的核心逻辑是通过index_from_dataset提前将候选向量存入内部索引结构,后续调用时仅需传入查询向量,内部直接基于预构建的索引完成TopK检索。这种设计针对静态候选库优化,无法直接通过index(input_query, input_candidates, k=5)的方式动态传入候选集。

实现动态候选匹配的可行方案

如果需要每次调用时指定不同的候选集,有两种可行方式:

1. 手动实现动态相似度计算与TopK排序

直接基于模型的查询分支(query_model)和候选分支(candidate_model)计算向量,实时计算相似度并筛选TopK结果,无需依赖FactorizedTopK层。

实现模型的call方法

class RetrievalModel(tf.keras.Model):
    def __init__(self, query_model, candidate_model):
        super().__init__()
        self.query_model = query_model
        self.candidate_model = candidate_model

    def call(self, inputs):
        # inputs为包含查询样本和候选样本的元组
        input_query, input_candidates = inputs

        # 生成查询向量和候选向量
        query_embeds = self.query_model(input_query)
        candidate_embeds = self.candidate_model(input_candidates)

        # 计算相似度(这里用内积,需与训练时的损失函数匹配)
        # 输出形状:(查询批次大小, 候选样本数量)
        similarities = tf.linalg.matmul(query_embeds, candidate_embeds, transpose_b=True)

        # 取TopK的相似度值和对应索引
        top_k_vals, top_k_indices = tf.math.top_k(similarities, k=5)

        # 从候选集中提取对应ID(假设候选样本包含'id'字段)
        candidate_ids = input_candidates['id']
        top_k_ids = tf.gather(candidate_ids, top_k_indices)

        return top_k_ids, top_k_vals

使用示例

# 实例化模型
final_model = RetrievalModel(query_model=your_query_model, candidate_model=your_candidate_model)

# 准备输入:将候选集转换为模型可处理的张量(小候选集适用)
# 若候选集较大,可分批次计算后合并结果
input_candidates = next(iter(parsed_topK.batch(len(parsed_topK))))
# 输入查询样本和候选样本
top_k_ids, top_k_scores = final_model((input_query, input_candidates))

2. 针对超大动态候选集的优化

如果候选集规模极大,实时计算相似度效率过低,可考虑:

  • 每次调用前临时构建BruteForce/ScaNN索引(但会增加额外的索引构建开销)
  • 改用近似最近邻检索库,实时传入候选向量构建临时索引后检索

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 15:02:29