导入Keras训练的TensorFlow模型报Functional无predict_segmentation属性错
问题根源
你的问题既不属于保存环节的bug,也不属于加载环节的bug,核心原因是:vgg_unet是keras-segmentation库二次封装的模型类实例,它自带的train、predict_segmentation都是库额外封装的上层方法,不属于Keras原生模型的内置属性。你用原生tf.keras.models.save_model保存模型时,只会存储模型的计算图结构和权重参数,不会保存第三方库附加的自定义方法,因此用tf.keras.models.load_model加载出来的只是Keras原生的Functional基础模型对象,自然没有predict_segmentation方法。
注:所有方案都需要保证初始化模型时的类别数、输入尺寸等参数和训练时完全一致,否则会出现权重维度不匹配的报错。
解决方法
下面是两种常用的修复方案:
方案1:使用keras-segmentation自带的检查点加载
你训练时已经配置了checkpoints_path参数,训练过程中库已经自动生成了带自定义方法的检查点文件,直接按如下方式加载即可:
from keras_segmentation.models.unet import vgg_unet # 初始化和训练时结构完全一致的模型 model = vgg_unet(n_classes=50, input_height=512, input_width=608) # 加载训练好的检查点权重,epoch为5的话路径对应为/tmp/vgg_unet_1.00005,可自行核对文件后缀 model.load_weights("/tmp/vgg_unet_1.00005") # 此时就可以正常调用predict_segmentation方法 out = model.predict_segmentation( inp=image_to_test, out_fname="/tmp/out.png" )
方案2:复用已保存的hdf5权重文件
如果你要使用之前导出的my_model.hdf5文件,也可以将权重赋值给新初始化的keras-segmentation模型实例:
from keras_segmentation.models.unet import vgg_unet # 初始化对应结构的封装模型 model = vgg_unet(n_classes=50, input_height=512, input_width=608) # 加载你之前保存的hdf5权重 model.load_weights("my_model.hdf5") # 即可正常调用自定义方法 out = model.predict_segmentation( inp=image_to_test, out_fname="/tmp/out.png" )
内容的提问来源于stack exchange,提问作者Tonino
相关产品推荐
相关产品推荐

