如何在树莓派上运行tf.lite模型替代已保存的Keras模型
错误原因
tensorflow.keras.models.load_model()接口仅支持加载.h5格式的Keras模型、以及SavedModel文件夹格式的模型,无法直接读取.tflite格式文件。你传入.tflite文件路径后,接口会默认把该路径当成SavedModel文件夹,尝试读取文件夹内的saved_model.pb文件,自然会抛出路径不存在的错误。
修复方案
你需要替换模型加载逻辑和预测逻辑,适配TFLite模型的推理流程:
- 替换模型加载代码
将原代码中模型加载部分:
model = tensorflow.keras.models.load_model("yourmodel.tflite")
替换为TFLite专属的加载逻辑:
# 加载TFLite模型 interpreter = tensorflow.lite.Interpreter(model_path="yourmodel.tflite") interpreter.allocate_tensors() # 获取输入输出张量的索引,后续推理需要用到 input_details = interpreter.get_input_details() output_details = interpreter.get_output_details()
如果你的树莓派上单独安装了轻量版tflite_runtime,也可以用更小开销的导入方式:
import tflite_runtime.interpreter as tflite interpreter = tflite.Interpreter(model_path="yourmodel.tflite") interpreter.allocate_tensors() input_details = interpreter.get_input_details() output_details = interpreter.get_output_details()
- 替换预测逻辑代码
将原代码中预测部分的两行:
predictions = model.predict(img) classIndex = model.predict_classes(img)
替换为TFLite的推理逻辑:
# 给输入张量赋值 interpreter.set_tensor(input_details[0]['index'], img.astype('float32')) # 执行推理 interpreter.invoke() # 读取输出结果 predictions = interpreter.get_tensor(output_details[0]['index']) classIndex = np.argmax(predictions)
额外注意事项
- 转换TFLite模型时要确保输入输出的形状、数据类型和原Keras模型一致,你当前代码中已经把输入处理成了
(1,32,32,1)的归一化float格式,只要转换模型时没有修改输入规格,就可以正常运行。 - TFLite推理的CPU占用会远低于原Keras模型,非常适合树莓派这类边缘设备运行。
内容的提问来源于stack exchange,提问作者Aydn Altun
相关产品推荐
相关产品推荐

