如何在Python中从文件加载TensorFlow Lite模型
加载TensorFlow Lite Model Maker生成的图像分类模型
你想要的是加载后能直接使用Model Maker的evaluate、predict_top_k等API的模型对象,而非单纯的TFLite解释器。正确的加载方式是调用image_classifier.ImageClassifier.load()方法,而非你尝试的image_classifier.Load()(注意方法名的大小写和格式)。
完整加载代码
from tflite_model_maker import image_classifier from tflite_model_maker.image_classifier import DataLoader # 加载模型,使用ImageClassifier类的load方法 model = image_classifier.ImageClassifier.load('/path/to/model.tflite')
加载后执行测试代码
加载完成后,即可运行你需要的测试逻辑,注意修正原代码中的变量错误(比如test_data = data需改为test_data = test):
test = DataLoader.from_folder('/path/to/testImages') loss, accuracy = model.evaluate(test) # 辅助函数:根据两个输入是否匹配返回'black'/'red' def get_label_color(val1, val2): if val1 == val2: return 'black' else: return 'red' # 绘制100张测试图及预测标签,错误预测用红色标注 test_data = test import matplotlib.pyplot as plt plt.figure(figsize=(20, 20)) predicts = model.predict_top_k(test_data) for i, (image, label) in enumerate(test_data.gen_dataset().unbatch().take(100)): ax = plt.subplot(10, 10, i+1) plt.xticks([]) plt.yticks([]) plt.grid(False) plt.imshow(image.numpy(), cmap=plt.cm.gray) predict_label = predicts[i][0][0] color = get_label_color(predict_label, test_data.index_to_label[label.numpy()]) ax.xaxis.label.set_color(color) plt.xlabel(f'Predicted: {predict_label}') plt.show()
说明
你之前查看的普通TFLite加载文档,是针对直接用Interpreter运行推理的场景。而Model Maker生成的模型,通过其内置的ImageClassifier.load()方法加载后,会成为封装好的对象,可直接使用Model Maker提供的高级API,无需自行处理输入输出的格式转换。
内容的提问来源于stack exchange,提问作者scryptonight
相关产品推荐
相关产品推荐

