SageMaker中TFRS初始化FactorizedTopK报错:无法转换'counter'为形状
问题概述
使用TensorFlow 2.13.0和TensorFlow Recommenders(TFRS)构建推荐系统,在AWS SageMaker环境中初始化RecommendationModel内的FactorizedTopK指标时触发以下错误:
ValueError: Cannot convert '('c', 'o', 'u', 'n', 't', 'e', 'r')' to a shape. Found invalid entry 'c' of type '<class 'str'>'
错误发生在tfrs.metrics.FactorizedTopK的Streaming层添加名为"counter"的权重时,且仅在SageMaker环境(CPU/GPU实例均存在)出现,本地或Google Colab环境无此问题。
错误原因分析
从报错栈可定位核心问题:
TFRS的Streaming层初始化时调用self.add_weight("counter", dtype=tf.int32, trainable=False),但TensorFlow 2.13中Keras的add_weight方法参数签名已变更——第一个参数为shape而非name。这导致字符串"counter"被误当作shape参数传入,系统尝试将字符串拆分为字符元组作为形状,最终触发类型错误。
该差异仅在SageMaker环境出现,本质是环境中TFRS版本与TensorFlow 2.13不兼容:旧版TFRS仍沿用Keras旧版add_weight的参数顺序(第一个参数为name),而TensorFlow 2.13的Keras已调整参数顺序。
解决方案建议
1. 升级TFRS到兼容版本
在SageMaker环境中安装与TensorFlow 2.13匹配的TFRS版本,最新版已修复该参数传递问题:
pip install --upgrade tensorflow-recommenders
验证安装版本:
import tensorflow_recommenders as tfrs print(tfrs.__version__)
建议安装0.7.3及以上版本(该版本已适配TF 2.13的Keras API变化)。
2. 手动修正TFRS代码(临时方案)
若无法升级TFRS,可直接修改SageMaker环境中TFRS的Streaming层代码:
找到tensorflow_recommenders/layers/factorized_top_k.py文件,定位到Streaming.__init__方法中的add_weight调用:
# 原代码 self._counter = self.add_weight("counter", dtype=tf.int32, trainable=False)
修改为显式指定参数名的调用:
# 修改后 self._counter = self.add_weight(name="counter", dtype=tf.int32, trainable=False, shape=())
通过显式指定参数名,避免因参数顺序导致的误解。
3. 调整模型初始化逻辑
另一种临时方案是提前构建候选索引,避免在FactorizedTopK初始化时动态创建Streaming层:
# 提前构建候选嵌入数据集 candidate_embeddings = programs.batch(128).map(self.candidate_model) # 显式创建Streaming索引 streaming_index = tfrs.layers.factorized_top_k.Streaming(k=10).index_from_dataset(candidate_embeddings) # 初始化任务时传入已构建的索引 self.task = tfrs.tasks.Retrieval( metrics=tfrs.metrics.FactorizedTopK(candidates=streaming_index) )
这种方式绕过了FactorizedTopK内部自动创建Streaming层的逻辑,避免参数传递错误。
代码验证
修改完成后,重新初始化模型并测试:
model = RecommendationModel([128, 64, 32]) model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.0001))
若不再触发ValueError,则说明问题已解决。
内容的提问来源于stack exchange,提问作者Sulove Dahal

