使用TensorFlow Recommenders构建检索索引时触发AttributeError错误
使用TensorFlow Recommenders BruteForce组件时的数据集格式问题解决
问题场景
使用TensorFlow Recommenders的BruteForce构建检索索引:
index = tfrs.layers.factorized_top_k.BruteForce(model.customer_model, k = 400)
候选数据集通过tf.data.Dataset.zip构建,结构为:
<ZipDataset element_spec=({'article_id': TensorSpec(shape=(None,), dtype=tf.string, name=None), 'prod_name': TensorSpec(shape=(None,), dtype=tf.string, name=None), 'product_type_name': TensorSpec(shape=(None,), dtype=tf.string, name=None)}, TensorSpec(shape=(None, 64), dtype=tf.float32, name=None))>
执行索引构建代码时:
index.index_from_dataset(candidates)
触发如下错误:
AttributeError Traceback (most recent call last) Input In [28], in <cell line: 6>() 4 candidates = tf.data.Dataset.zip((articles.batch(128), articles.batch(128).map(model.article_model))) 5 print(candidates) ----> 6 index.index_from_dataset(candidates) File ~/miniconda3/envs/tf/lib/python3.9/site-packages/tensorflow_recommenders/layers/factorized_top_k.py:197, in TopK.index_from_dataset(self, candidates) 174 def index_from_dataset( 175 self, 176 candidates: tf.data.Dataset 177 ) -> "TopK": 178 """Builds the retrieval index. 179 180 When called multiple times the existing index will be dropped and a new one (...) 194 ValueError if the dataset does not have the correct structure. 195 """ --> 197 _check_candidates_with_identifiers(candidates) 199 spec = candidates.element_spec 201 if isinstance(spec, tuple): File ~/miniconda3/envs/tf/lib/python3.9/site-packages/tensorflow_recommenders/layers/factorized_top_k.py:127, in _check_candidates_with_identifiers(candidates) 119 raise ValueError( 120 "The dataset must yield candidate embeddings or " 121 "tuples of (candidate identifiers, candidate embeddings). " 122 f"Got {spec} instead." 123 ) 125 identifiers_spec, candidates_spec = spec --> 127 if candidates_spec.shape[0] != identifiers_spec.shape[0]: 128 raise ValueError( 129 "Candidates and identifiers have to have the same batch dimension. " 130 f"Got {candidates_spec.shape[0]} and {identifiers_spec.shape[0]}." 131 ) AttributeError: 'dict' object has no attribute 'shape'
错误原因
错误根源是候选数据集的标识符部分是字典类型,而TensorFlow Recommenders内部的_check_candidates_with_identifiers函数期望标识符是单个Tensor(带有shape属性)。字典没有shape属性,导致维度检查时触发AttributeError。
解决方案
BruteForce组件只需要唯一的候选标识符(比如article_id)来返回检索结果,不需要传入所有特征字典。只需修改候选数据集的构建逻辑,提取单个标识符Tensor即可:
# 提取唯一标识符article_id作为候选标识 candidate_ids = articles.batch(128).map(lambda x: x["article_id"]) # 生成候选嵌入 candidate_embeddings = articles.batch(128).map(model.article_model) # 构建符合要求的候选数据集 candidates = tf.data.Dataset.zip((candidate_ids, candidate_embeddings)) # 构建索引 index.index_from_dataset(candidates)
修改后,标识符部分变为单个Tensor,具备shape属性,能够通过内部的维度检查,顺利完成索引构建。
内容的提问来源于stack exchange,提问作者Claudiu Stoica
相关产品推荐
相关产品推荐

