TensorFlow模型转TFLite时签名键相关错误排查求助
TensorFlow语音识别模型转TFLite报错解决
问题背景
基于TensorFlow官方语音识别教程的Notebook,使用自定义数据集训练并导出模型后,转换为TFLite格式时出现错误,使用TensorFlow 2.11.0版本。
首次转换尝试与报错
转换代码:
# Load the saved model saved_model_path = "saved" saved_model = tf.saved_model.load(saved_model_path) # Set the concrete function to be used for conversion concrete_func = saved_model.signatures['serving_default'] # Convert the model to TFLite and save it in a new folder called "saved-lite" converter = tf.lite.TFLiteConverter.from_concrete_functions([concrete_func]) tflite_model = converter.convert() with open('saved-lite/model.tflite', 'wb') as f: f.write(tflite_model)
报错信息:
KeyError: 'serving_default'
错误栈:
--------------------------------------------------------------------------- KeyError Traceback (most recent call last) ~\AppData\Local\Temp\ipykernel_28532\3665239774.py in <module> 4 5 # Set the concrete function to be used for conversion ----> 6 concrete_func = saved_model.signatures['serving_default'] 7 8 # Convert the model to TFLite and save it in a new folder called "saved-lite" ~\AppData\Roaming\Python\Python39\site-packages\tensorflow\python\saved_model\signature_serialization.py in __getitem__(self, key) 245 246 def __getitem__(self, key): --> 247 return self._signatures[key] 248 249 def __iter__(self):
二次转换尝试与报错
转换代码:
# Converting a SavedModel to a TensorFlow Lite model. converter = tf.lite.TFLiteConverter.from_saved_model('saved') tflite_model = converter.convert()
报错信息:
ValueError: Only support at least one signature key.
导出模型代码
使用教程提供的代码导出模型:
class ExportModel(tf.Module): def __init__(self, model): self.model = model # Accept either a string-filename or a batch of waveforms. # You could add additional signatures for a single wave, or a ragged-batch. self.__call__.get_concrete_function( x=tf.TensorSpec(shape=(), dtype=tf.string)) self.__call__.get_concrete_function( x=tf.TensorSpec(shape=[None, 16000], dtype=tf.float32)) @tf.function def __call__(self, x): # If they pass a string, load the file and decode it. if x.dtype == tf.string: x = tf.io.read_file(x) x, _ = tf.audio.decode_wav(x, desired_channels=1, desired_samples=16000,) x = tf.squeeze(x, axis=-1) x = x[tf.newaxis, :] x = get_spectrogram(x) result = self.model(x, training=False) class_ids = tf.argmax(result, axis=-1) class_names = tf.gather(label_names, class_ids) return {'predictions':result, 'class_ids': class_ids, 'class_names': class_names} export = ExportModel(model)
模型已保存至saved/目录,目录结构如图:
问题原因与解决方法
问题出在模型导出时未显式为签名设置名称并注册到SavedModel,导致转换时找不到默认的serving_default签名,且TFLite转换器无法识别有效签名。
修正后的导出与转换步骤
- 修改模型导出代码:为每个输入签名命名,并在保存时显式指定签名映射:
class ExportModel(tf.Module): def __init__(self, model): self.model = model # 为两种输入分别创建带名称的签名 self.file_input = self.__call__.get_concrete_function( x=tf.TensorSpec(shape=(), dtype=tf.string)) self.waveform_input = self.__call__.get_concrete_function( x=tf.TensorSpec(shape=[None, 16000], dtype=tf.float32)) @tf.function def __call__(self, x): # 原有逻辑不变 if x.dtype == tf.string: x = tf.io.read_file(x) x, _ = tf.audio.decode_wav(x, desired_channels=1, desired_samples=16000,) x = tf.squeeze(x, axis=-1) x = x[tf.newaxis, :] x = get_spectrogram(x) result = self.model(x, training=False) class_ids = tf.argmax(result, axis=-1) class_names = tf.gather(label_names, class_ids) return {'predictions':result, 'class_ids': class_ids, 'class_names': class_names} export = ExportModel(model) # 保存模型时显式指定签名映射 tf.saved_model.save(export, "saved", signatures={ 'serving_default': export.file_input, # 设置默认签名为文件输入,也可选择waveform_input 'waveform_input': export.waveform_input })
- 重新执行TFLite转换:
converter = tf.lite.TFLiteConverter.from_saved_model('saved') tflite_model = converter.convert() # 保存TFLite模型 with open('saved-lite/model.tflite', 'wb') as f: f.write(tflite_model)
另一种直接转换方式(无需重新导出模型)
如果不想重新导出模型,可加载模型后获取已存在的具体函数进行转换:
saved_model = tf.saved_model.load("saved") # 获取所有可用的具体函数 concrete_funcs = list(saved_model.signatures.values()) # 取第一个具体函数进行转换 converter = tf.lite.TFLiteConverter.from_concrete_functions([concrete_funcs[0]]) tflite_model = converter.convert() with open('saved-lite/model.tflite', 'wb') as f: f.write(tflite_model)
内容的提问来源于stack exchange,提问作者Yasiru Ruwantha Weerakoon
相关产品推荐
相关产品推荐

