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

TensorFlow Recommender排序模型加载后预测报错问题求助

TensorFlow Recommender排序模型保存加载后预测报错问题

问题描述

用TensorFlow Recommender构建排序模型,原模型直接预测运行正常,但保存后重新加载再执行预测时出现报错。移除年龄相关数据及模型层后,加载预测恢复正常,推测问题出在年龄相关模块,但无法定位具体原因。

排序模型代码

class HMRankingModel(tfrs.Model):

    def __init__(self):
        super().__init__()
        
        # 用户塔
        self.customer_input = tf.keras.Input(shape=(1,), dtype=tf.string, name='customer_input')
        self.customer_sl = tf.keras.layers.StringLookup(vocabulary=unique_customer_ids, mask_token=None, name='customer_string_lookup')(self.customer_input)
        self.customer_embedding = tf.squeeze(tf.keras.layers.Embedding(len(unique_customer_ids) + 1, embedding_dimension, name='customer_emb')(self.customer_sl), axis=1)

        self.age_input = tf.keras.Input(shape=(1,), name='age_input')
        self.age_discretization = tf.keras.layers.Discretization(age_buckets.tolist(), name='age_discretization')(self.age_input)
        self.age_embedding = tf.squeeze(tf.keras.layers.Embedding(len(age_buckets) + 1, embedding_dimension, name='age_embedding')(self.age_discretization), axis=1)
        
        self.customer_merged = tf.keras.layers.concatenate([self.customer_embedding, self.age_embedding], axis=-1, name='customer_merged')
        self.customer_dense = tf.keras.layers.Dense(embedding_dimension, activation=activation, name='customer_dense')(self.customer_merged)
        
        
        # 物品塔
        self.article_input = tf.keras.Input(shape=(1,), dtype=tf.string, name='article_input')
        self.article_sl = tf.keras.layers.StringLookup(vocabulary=unique_article_ids, name='article_string_lookup')(self.article_input)
        self.article_final = tf.squeeze(tf.keras.layers.Embedding(len(unique_article_ids)+1, embedding_dimension, name='article_emb')(self.article_sl), axis=1)

        self.article_dense = tf.keras.layers.Dense(embedding_dimension, activation=activation, name='article_dense')(self.article_final)        


        # 交互层
        self.towers_multiplied = tf.keras.layers.Multiply(name='towers_multiplied')([self.customer_dense, self.article_dense])
        self.towers_dense = tf.keras.layers.Dense(dense_size, activation=activation, name='towers_dense1')(self.towers_multiplied)
        self.output_node = tf.keras.layers.Dense(1, name='output_node')(self.towers_dense)
        
        
        # 模型定义
        self.model = tf.keras.Model(inputs={'customer_id': self.customer_input, 
                                            'article_id': self.article_input,
                                            'age': self.age_input,
                                            }, 
                                    outputs=self.output_node)
        
        self.task = tfrs.tasks.Ranking(
            loss = tf.keras.losses.MeanSquaredError(),
            metrics=[tf.keras.metrics.RootMeanSquaredError()]
        )
        
    def call(self, features):
        return self.model({'customer_id': features["customer_id"], 
                           'article_id': features["article_id"], 
                           'age': features["age"], 
                           })   
        
    def compute_loss(self, features_dict, training=False):
        labels = features_dict.pop("count")
        predictions = self(features_dict)
        return self.task(labels=labels, predictions=predictions)

模型训练代码

ranking_model = HMRankingModel()
ranking_model.compile(optimizer=tf.keras.optimizers.Adagrad(learning_rate=0.1))
ranking_model.fit(cached_train, validation_data=cached_validation, epochs=epochs)

原模型预测(正常运行)

ranking_model({
    'customer_id': np.array(["18b3a4767533c8f1f6ff274b57ca200939c9fda3992c5bb3b50b31dc6d6b1ee5"]), 
    'age': np.array([29]), 
    'article_id': np.array(['562245059'])
})

输出:

<tf.Tensor: shape=(1, 1), dtype=float32, numpy=array([[1.3872527]], dtype=float32)>

模型保存与加载预测(报错)

tf.saved_model.save(ranking_model, ranking_model_path)

saved_ranking_model = tf.saved_model.load(ranking_model_path)
predictions = saved_ranking_model({
    'customer_id': np.array(["18b3a4767533c8f1f6ff274b57ca200939c9fda3992c5bb3b50b31dc6d6b1ee5"]), 
    'age': np.array([29]), 
    'article_id': np.array(['141661025'])
})

报错信息:

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
Input In [31], in <cell line: 1>()
----> 1 predictions = saved_ranking_model({
      2     'customer_id': np.array(["18b3a4767533c8f1f6ff274b57ca200939c9fda3992c5bb3b50b31dc6d6b1ee5"]), 
      3     'age': np.array([29]), 
      4     'article_id': np.array(['141661025'])
      5 })

File ~/miniconda3/envs/tf/lib/python3.9/site-packages/tensorflow/python/saved_model/load.py:686, in _call_attribute(instance, *args, **kwargs)
    685 def _call_attribute(instance, *args, **kwargs):
--> 686   return instance.__call__(*args, **kwargs)

File ~/miniconda3/envs/tf/lib/python3.9/site-packages/tensorflow/python/util/traceback_utils.py:153, in filter_traceback.<locals>.error_handler(*args, **kwargs)
    151 except Exception as e:
    152   filtered_tb = _process_traceback_frames(e.__traceback__)
--> 153   raise e.with_traceback(filtered_tb) from None
    154 finally:
    155   del filtered_tb

File ~/miniconda3/envs/tf/lib/python3.9/site-packages/tensorflow/python/saved_model/function_deserialization.py:286, in recreate_function.<locals>.restored_function_body(*args, **kwargs)
    282   positional, keyword = concrete_function.structured_input_signature
    283   signature_descriptions.append(
    284       "Option {}:
  {}
  Keyword arguments: {}"
    285       .format(index + 1, _pretty_format_positional(positional), keyword))
--> 286 raise ValueError(
    287     "Could not find matching concrete function to call loaded from the "
    288     f"SavedModel. Got:
  {_pretty_format_positional(args)}
  Keyword "
    289     f"arguments: {kwargs}

 Expected these arguments to match one of the "
    290     f"following {len(saved_function.concrete_functions)} option(s):

"
    291     f"{(chr(10)+chr(10)).join(signature_descriptions)}")

ValueError: Could not find matching concrete function to call loaded from the SavedModel. Got:
  Positional arguments (2 total):
    * {'age': <tf.Tensor 'features:0' shape=(1,) dtype=int64>,
 'article_id': <tf.Tensor 'features_1:0' shape=(1,) dtype=string>,
 'customer_id': <tf.Tensor 'features_2:0' shape=(1,) dtype=string>}
    * False
  Keyword arguments: {}

 Expected these arguments to match one of the following 4 option(s):

Option 1:
  Positional arguments (2 total):
    * {'age': TensorSpec(shape=(None,), dtype=tf.float32, name='age'),
 'article_id': TensorSpec(shape=(None,), dtype=tf.string, name='article_id'),
 'customer_id': TensorSpec(shape=(None,), dtype=tf.string, name='customer_id')}
    * False
  Keyword arguments: {}

Option 2:
  Positional arguments (2 total):
    * {'age': TensorSpec(shape=(None,), dtype=tf.float32, name='features/age'),
 'article_id': TensorSpec(shape=(None,), dtype=tf.string, name='features/article_id'),
 'customer_id': TensorSpec(shape=(None,), dtype=tf.string, name='features/customer_id')}
    * False
  Keyword arguments: {}

Option 3:
  Positional arguments (2 total):
    * {'age': TensorSpec(shape=(None,), dtype=tf.float32, name='features/age'),
 'article_id': TensorSpec(shape=(None,), dtype=tf.string, name='features/article_id'),
 'customer_id': TensorSpec(shape=(None,), dtype=tf.string, name='features/customer_id')}
    * True
  Keyword arguments: {}

Option 4:
  Positional arguments (2 total):
    * {'age': TensorSpec(shape=(None,), dtype=tf.float32, name='age'),
 'article_id': TensorSpec(shape=(None,), dtype=tf.string, name='article_id'),
 'customer_id': TensorSpec(shape=(None,), dtype=tf.string, name='customer_id')}
    * True
  Keyword arguments: {}

问题原因

从报错信息能直接看到:输入的age数据类型是int64,但加载后的模型期望输入类型为float32。原模型预测时TensorFlow会自动做类型转换,而保存加载后的模型对输入类型检查更严格,类型不匹配直接触发报错。另外,模型中age_input定义时未显式指定数据类型,默认是float32,和输入的int64数组不兼容。

解决方案

方案1:修改输入数据类型

预测时显式将age数组转为float32类型:

predictions = saved_ranking_model({
    'customer_id': np.array(["18b3a4767533c8f1f6ff274b57ca200939c9fda3992c5bb3b50b31dc6d6b1ee5"]), 
    'age': np.array([29], dtype=np.float32),  # 指定float32类型
    'article_id': np.array(['141661025'])
})

方案2:显式指定age_input的数据类型

在模型定义时,为age_input指定dtype=tf.int64,匹配输入的整数类型:

self.age_input = tf.keras.Input(shape=(1,), dtype=tf.int64, name='age_input')

方案3:在call方法中添加类型转换

在模型的call方法里对age特征做类型转换,确保和模型层输入要求一致:

def call(self, features):
    return self.model({
        'customer_id': features["customer_id"], 
        'article_id': features["article_id"], 
        'age': tf.cast(features["age"], tf.float32),  # 转换为float32
    })   

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 00:40:42