tf.keras保存的H5模型无法在Keras 2.2.4加载,求转换方案
解决tf.keras保存的H5模型无法在Keras 2.2.4加载的问题
你的问题核心在于tf.keras与早期原生Keras(2.2.4版本)的API差异:tf.keras为了支持TensorFlow特有的功能(比如ragged张量),在部分层(比如Embedding)的配置中加入了原生Keras 2.2.4不支持的ragged参数,导致加载时触发参数不匹配的错误。下面提供几种可行的解决方法:
方法一:修改模型配置后重新保存(推荐)
在有TensorFlow的环境中加载模型,移除配置里不兼容的参数,再保存为原生Keras可读取的版本:
- 首先用tf.keras加载原模型:
import tensorflow as tf # 加载原模型 model = tf.keras.models.load_model("/home/Documents/explorePrj/Segmentation/models/model.h5", compile=False)
- 遍历所有层,清理不兼容的参数:
def clean_layer_config(layer_config): # 移除原生Keras不支持的ragged参数 if 'ragged' in layer_config['config']: del layer_config['config']['ragged'] # 如果还有其他不兼容参数,也可以在这里添加删除逻辑 return layer_config # 获取模型配置并清理所有层的配置 model_config = model.get_config() model_config['layers'] = [clean_layer_config(layer) for layer in model_config['layers']] # 用清理后的配置重建模型,并复制原权重 fixed_model = tf.keras.Model.from_config(model_config) fixed_model.set_weights(model.get_weights())
- 保存修改后的模型:
fixed_model.save("/home/Documents/explorePrj/Segmentation/models/compatible_model.h5", save_format='h5')
现在这个新保存的模型应该可以在Keras 2.2.4中正常加载了。
方法二:自定义兼容层,通过custom_objects加载
如果不想重新保存模型,可以在Keras 2.2.4环境中,自定义一个忽略不兼容参数的层类,加载时指定custom_objects:
以报错涉及的Embedding层为例:
from keras.layers import Embedding from keras.legacy import interfaces class CompatibleEmbedding(Embedding): @interfaces.legacy_embedding_support def __init__(self, input_dim, output_dim, embeddings_initializer='uniform', embeddings_regularizer=None, activity_regularizer=None, embeddings_constraint=None, mask_zero=False, input_length=None, ragged=False, **kwargs): # 直接忽略ragged参数,调用父类初始化 super().__init__(input_dim, output_dim, embeddings_initializer=embeddings_initializer, embeddings_regularizer=embeddings_regularizer, activity_regularizer=activity_regularizer, embeddings_constraint=embeddings_constraint, mask_zero=mask_zero, input_length=input_length, **kwargs) # 加载模型时指定自定义层 from keras.models import load_model model = load_model("/home/Documents/explorePrj/Segmentation/models/model.h5", compile=False, custom_objects={'Embedding': CompatibleEmbedding})
如果还有其他层出现类似参数错误,需要对应编写兼容类并添加到custom_objects中。
方法三:升级原生Keras(备选)
如果你的环境允许,将原生Keras升级到2.3.x及以上版本——从Keras 2.3开始,官方已经完成了与tf.keras的API对齐,大部分参数差异被消除,直接加载tf.keras保存的模型应该不会再出现此类问题。
内容的提问来源于stack exchange,提问作者ShivSd
相关产品推荐
相关产品推荐

