TensorFlow 2中MetaGraph转SavedModel遇Trackable类型错误求助
解决ConcreteFunction无法保存为SavedModel的问题
核心原因
报错本质是tf.saved_model.save仅接受继承自Trackable的对象(比如tf.Module),而你得到的WrappedFunction/ConcreteFunction不属于这类对象,直接保存会触发类型校验错误。
解决方法:将ConcreteFunction包装进tf.Module
把推理函数封装到简单的tf.Module子类中,使其符合SavedModel的导出要求,代码示例如下:
import tensorflow as tf # 假设pruned_func是你通过wrapped_import.prune得到的函数 pruned_func = ... # 定义包装用的Module类 class ModelWrapper(tf.Module): def __init__(self, func): super().__init__() self.inference_func = func @tf.function(input_signature=[tf.TensorSpec(shape=[None, 224, 224, 3], dtype=tf.float32, name='input_image')]) def __call__(self, input_image): return self.inference_func(input_image) # 完成函数包装 wrapped_model = ModelWrapper(pruned_func) # 保存为SavedModel tf.saved_model.save(wrapped_model, "./model_files")
关键说明
- 必须指定
input_signature:这是导出SavedModel的必要条件,要替换成你模型实际的输入形状、数据类型和名称。 - 包装后的
ModelWrapper属于Trackable子类,满足tf.saved_model.save的要求。
后续转换TF-TRT模型
SavedModel保存成功后,即可用TF-TRT API完成转换,示例代码:
from tensorflow.python.compiler.tensorrt import trt_convert as trt converter = trt.TrtGraphConverterV2(input_saved_model_dir="./model_files") converter.convert() converter.save("./trt_model")
内容的提问来源于stack exchange,提问作者Numan988
相关产品推荐
相关产品推荐

