TensorFlow升级至2.7以上版本后tf.keras.load_model无法加载旧模型求助
解决TensorFlow 2.7保存模型在高版本加载的 dtype 不兼容问题
问题核心
你遇到的是TF版本升级后,模型保存时的输入签名 dtype(float16)与高版本TF加载时默认的 dtype(float32)不匹配,导致类型校验失败的问题。以下是无需修改TF源码的规避方案:
方案1:在TF2.7环境中转存模型,统一输入 dtype
利用TF2.7能正常加载模型的特性,重新定义输入层的 dtype 后保存新模型,让高版本TF可以正常识别:
import tensorflow as tf # 确认当前环境为TensorFlow 2.7.x print(tf.__version__) # 加载原模型 original_model = tf.keras.models.load_model("path/to/your/original/model") # 重新定义输入层,指定统一的dtype(这里用float32,也可根据原模型选择float16) input_layer = tf.keras.layers.Input( shape=original_model.input_shape[1:], # 保留除batch外的维度 dtype=tf.float32, name=original_model.input_names[0] ) # 复用原模型的层结构 x = input_layer for layer in original_model.layers[1:]: x = layer(x) # 构建新模型并复制原权重 new_model = tf.keras.Model(inputs=input_layer, outputs=x) new_model.set_weights(original_model.get_weights()) # 保存兼容高版本的新模型 new_model.save("path/to/compatible/model")
方案2:在高版本TF中修改加载逻辑,适配 dtype
直接在高版本环境中加载SavedModel,通过包装预测函数统一 dtype 后重新保存:
import tensorflow as tf # 加载原SavedModel模型 loaded_model = tf.saved_model.load("path/to/your/original/model") # 定义包装函数,统一输入输出的dtype @tf.function(input_signature=[tf.TensorSpec(shape=(None, None, 80), dtype=tf.float32, name='x')]) def wrapped_infer(inputs): # 将输入转换为原模型预期的float16 inputs_cast = tf.cast(inputs, tf.float16) outputs = loaded_model(inputs_cast) # 可选:将输出转换为float32适配高版本默认设置 return tf.cast(outputs, tf.float32) # 保存修改后的模型 tf.saved_model.save(wrapped_infer, "path/to/fixed/model") # 现在可以用Keras正常加载 compatible_model = tf.keras.models.load_model("path/to/fixed/model")
方案3:分离模型结构与权重,重新加载
放弃完整模型保存的元数据,只保存权重,在高版本中重新构建模型结构后加载权重:
- 在TF2.7中导出权重
original_model = tf.keras.models.load_model("path/to/original/model") original_model.save_weights("model_weights.h5")
- 在高版本TF中重建模型并加载权重
import tensorflow as tf # 严格复刻原模型的结构(输入dtype可统一设为float32或float16) def build_original_model(): inputs = tf.keras.layers.Input(shape=(None, None, 80), dtype=tf.float32) # 替换为你原模型的所有层结构,确保与原模型完全一致 x = tf.keras.layers.Conv1D(64, 3, activation='relu')(inputs) outputs = tf.keras.layers.Dense(1)(x) return tf.keras.Model(inputs=inputs, outputs=outputs) # 构建模型并加载权重 model = build_original_model() model.load_weights("model_weights.h5")
内容的提问来源于stack exchange,提问作者Error404
相关产品推荐
相关产品推荐

