加载TensorFlow模型后,如何调用带自定义参数的call方法?
问题背景
我正在学习TensorFlow文本生成教程,其中包含MyModel和OneStep两个模型:MyModel是处理向量化字符串的RNN模型,OneStep则封装MyModel直接处理字符串。
教程中演示了OneStep模型的保存与加载,我已成功实现,但现在需要保存并重新加载MyModel。尝试调用加载后的模型并传入return_state=True时出现报错:
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) /tmp/ipykernel_23/2335414736.py in <module> 1 # TODO: Loaded model gives an error 2 for input_example_batch, target_example_batch in train_ds.take(1): ----> 3 example_batch_predictions, example_states = loaded_model(input_example_batch, False, None, return_state=True) 4 print(example_batch_predictions.shape, "# (batch_size, sequence_length, vocab_size)") 5 print(example_states.shape, " # (batch_size, rnn_units)") /opt/conda/lib/python3.7/site-packages/tensorflow/python/saved_model/load.py in _call_attribute(instance, *args, **kwargs) 662 663 def _call_attribute(instance, *args, **kwargs): --> 664 return instance.__call__(*args, **kwargs) 665 666 /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/def_function.py in __call__(self, *args, **kwds) 883 884 with OptionalXlaContext(self._jit_compile): --> 885 result = self._call(*args, **kwds) 886 887 new_tracing_count = self.experimental_get_tracing_count() /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/def_function.py in _call(self, *args, **kwds) 931 # This is the first call of __call__, so we have to initialize. 932 initializers = [] --> 933 self._initialize(args, kwds, add_initializers_to=initializers) 934 finally: 935 # At this point we know that the initialization is complete (or less /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/def_function.py in _initialize(self, args, kwds, add_initializers_to) 758 self._concrete_stateful_fn = ( 759 self._stateful_fn._get_concrete_function_internal_garbage_collected( # pylint: disable=protected-access --> 760 *args, **kwds)) 761 762 def invalid_creator_scope(*unused_args, **unused_kwds): /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/function.py in _get_concrete_function_internal_garbage_collected(self, *args, **kwargs) 3064 args, kwargs = None, None 3065 with self._lock: --> 3066 graph_function, _ = self._maybe_define_function(args, kwargs) 3067 return graph_function 3068 /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/function.py in _maybe_define_function(self, args, kwargs) 3461 3462 self._function_cache.missed.add(call_context_key) --> 3463 graph_function = self._create_graph_function(args, kwargs) 3464 self._function_cache.primary[cache_key] = graph_function 3465 /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/function.py in _create_graph_function(self, args, kwargs, override_flat_arg_shapes) 3306 arg_names=arg_names, 3307 override_flat_arg_shapes=override_flat_arg_shapes, --> 3308 capture_by_value=self._capture_by_value), 3309 self._function_attributes, 3310 function_spec=self.function_spec, /opt/conda/lib/python3.7/site-packages/tensorflow/python/framework/func_graph.py in func_graph_from_py_func(name, python_func, args, kwargs, signature, func_graph, autograph, autograph_options, add_control_dependencies, arg_names, op_return_value, collections, capture_by_value, override_flat_arg_shapes, acd_record_initial_resource_uses) 1005 _, original_func = tf_decorator.unwrap(python_func) 1006 --> 1007 func_outputs = python_func(*func_args, **func_kwargs) 1008 1009 # invariant: `func_outputs` contains only Tensors, CompositeTensors, /opt/conda/lib/python3.7/site-packages/tensorflow/python/eager/def_function.py in wrapped_fn(*args, **kwds) 666 # the function a weak reference to itself to avoid a reference cycle. 667 with OptionalXlaContext(compile_with_xla): --> 668 out = weak_wrapped_fn().__wrapped__(*args, **kwds) 669 return out 670 /opt/conda/lib/python3.7/site-packages/tensorflow/python/saved_model/function_deserialization.py in restored_function_body(*args, **kwargs) 292 .format(_pretty_format_positional(args), kwargs, 293 len(saved_function.concrete_functions), --> 294 "\n".join(signature_descriptions))) 295 296 concrete_function_objects = [] ValueError: Could not find matching function to call loaded from the SavedModel. Got: Positional arguments (4 total): * Tensor("inputs:0", shape=(64, 113), dtype=int64) * False * None * True Keyword arguments: {} Expected these arguments to match one of the following 4 option(s): Option 1: Positional arguments (4 total): * TensorSpec(shape=(None, 113), dtype=tf.int64, name='input_1') * False * None * False Keyword arguments: {} Option 2: Positional arguments (4 total): * TensorSpec(shape=(None, 113), dtype=tf.int64, name='inputs') * False * None * False Keyword arguments: {} Option 3: Positional arguments (4 total): * TensorSpec(shape=(None, 113), dtype=tf.int64, name='inputs') * True * None * False Keyword arguments: {} Option 4: Positional arguments (4 total): * TensorSpec(shape=(None, 113), dtype=tf.int64, name='input_1') * True * None * False Keyword arguments: {}
我认为问题出在call方法中的自定义参数,以下是复现该问题的最简示例:
import tensorflow as tf class CustomModel(tf.keras.models.Model): def __init__(self): super().__init__() self.dense = tf.keras.layers.Dense(10) def call(self, inputs, custom_param=False): return self.dense(inputs) model = CustomModel() sample_inputs = tf.zeros((16, 30)) print('Sample inputs:', sample_inputs) sample_outputs = model(sample_inputs) print('Sample outputs:', sample_outputs) model.save('saved_model') loaded_model = tf.keras.models.load_model('saved_model') sample_outputs_2 = loaded_model(sample_inputs, custom_param=True) print('Sample outputs 2:', sample_outputs_2)
调用加载后的模型时,只要custom_param使用非默认值就会失败。
请问这是Bug还是设计如此?如何修改模型,使其在训练时仅返回输出序列,推理时返回输出序列和状态,以便将状态回喂给模型生成更多字符?
解答
这是设计如此,不是Bug
TensorFlow SavedModel在保存时,只会记录模型实际被调用过的函数签名。在你的示例中,保存模型前只调用过model(sample_inputs)(使用custom_param=False的默认值),所以SavedModel中只保存了这个参数组合的函数签名。加载后调用非默认参数时,找不到匹配的签名,就会报错。
对于你的MyModel来说,保存前可能只在训练模式(return_state=False)下运行过,所以加载后无法直接调用return_state=True的版本。
修改方案:保存前触发所有需要的函数签名,或使用显式的函数签名定义
方案1:保存模型前,提前调用所有需要的参数组合
在调用model.save()之前,先调用一次带非默认参数的模型,让TensorFlow记录对应的函数签名:
# 保存前先触发一次自定义参数的调用 sample_outputs_custom = model(sample_inputs, custom_param=True) # 再保存模型 model.save('saved_model')
这样加载后,就能正常调用loaded_model(sample_inputs, custom_param=True)了。
对于你的RNN模型,就是在保存前先调用一次return_state=True的版本:
# 假设train_ds是你的训练数据集 for input_example_batch, _ in train_ds.take(1): # 触发一次return_state=True的调用 model(input_example_batch, False, None, return_state=True) # 再保存模型 model.save('my_model')
方案2:使用tf.function显式定义不同的调用签名(更优雅)
在模型类中定义不同的方法,用tf.function装饰并指定输入签名,这样保存时会记录这些签名:
import tensorflow as tf class CustomModel(tf.keras.models.Model): def __init__(self): super().__init__() self.dense = tf.keras.layers.Dense(10) def call(self, inputs, custom_param=False): return self.dense(inputs) # 定义显式的推理方法 @tf.function(input_signature=[tf.TensorSpec(shape=(None, 30), dtype=tf.float32)]) def infer_with_custom_param(self, inputs): return self.call(inputs, custom_param=True) # 使用示例 model = CustomModel() sample_inputs = tf.zeros((16, 30)) model(sample_inputs) # 训练模式调用 model.infer_with_custom_param(sample_inputs) # 推理模式调用 model.save('saved_model') # 加载后调用 loaded_model = tf.keras.models.load_model('saved_model') loaded_model.infer_with_custom_param(sample_inputs)
针对RNN返回状态的修改方案
对于你的MyModel,要实现训练返回输出、推理返回输出+状态,可以这样调整:
class MyModel(tf.keras.Model): def __init__(self, vocab_size, embedding_dim, rnn_units): super().__init__() self.embedding = tf.keras.layers.Embedding(vocab_size, embedding_dim) self.gru = tf.keras.layers.GRU(rnn_units, return_sequences=True, return_state=True) self.dense = tf.keras.layers.Dense(vocab_size) def call(self, inputs, states=None, return_state=False, training=False): x = inputs x = self.embedding(x, training=training) if states is None: states = self.gru.get_initial_state(x) x, states = self.gru(x, initial_state=states, training=training) x = self.dense(x, training=training) if return_state: return x, states else: return x # 显式定义推理用的方法,带return_state=True @tf.function(input_signature=[ tf.TensorSpec(shape=(None, None), dtype=tf.int64), tf.TensorSpec(shape=(None, None), dtype=tf.float32, name='states') ]) def generate(self, inputs, states): return self.call(inputs, states=states, return_state=True, training=False) # 保存前触发必要的签名 model = MyModel(vocab_size=..., embedding_dim=..., rnn_units=...) # 训练模式调用 for input_batch, target_batch in train_ds.take(1): model(input_batch) # 推理模式调用(触发return_state=True的签名) sample_states = tf.zeros((64, rnn_units)) model(input_batch, states=sample_states, return_state=True) # 保存模型 model.save('my_model') # 加载后使用 loaded_model = tf.keras.models.load_model('my_model') # 训练模式 output = loaded_model(input_batch) # 推理模式,返回输出和状态 output, new_states = loaded_model(input_batch, states=sample_states, return_state=True)
内容的提问来源于stack exchange,提问作者A Kubiesa

