如何通过TFLite Model Maker定制TFLite图像分类模型的输出
解决TFLite图像分类模型指定输出类别的方法
核心逻辑
TFLite Model Maker的图像分类模型,输出类别名称完全由训练数据集的文件夹命名决定,无需在训练时额外设置参数,关键是先整理好数据集的结构。
1. 按需求整理训练数据集的文件夹结构
- 将同一类别的图像放入同名文件夹,文件夹名称就是你想要的模型输出类别名称。比如你希望模型输出“cat”“dog”“bird”,就创建三个对应名称的文件夹,把各类图像分别放入。
- 示例结构:
training_data/ ├─ cat/ │ ├─ cat001.jpg │ ├─ cat002.jpg │ └─ ... ├─ dog/ │ ├─ dog001.jpg │ └─ ... └─ bird/ ├─ bird001.jpg └─ ... - 注意:文件夹名称尽量用无特殊字符的英文(避免编码问题),不要用系统自动生成的“image”“image1”这类名称。
2. 加载数据集并训练模型
直接指向整理好的数据集根目录加载数据,Model Maker会自动读取子文件夹名称作为类别标签。示例代码片段:
from tflite_model_maker import image_classifier from tflite_model_maker.image_classifier import DataLoader # 加载结构化数据集 data = DataLoader.from_folder('training_data/') train_data, test_data = data.split(0.8) # 训练模型 model = image_classifier.create(train_data) # 评估模型效果 loss, accuracy = model.evaluate(test_data)
3. 验证导出模型的类别标签
导出TFLite模型后,类别标签会嵌入模型中,可通过以下方式确认:
- 用Python读取模型标签:
from tflite_support.metadata import metadata_schema_py_generated as _metadata_fb from tflite_support.metadata import metadata as _metadata # 加载导出的TFLite模型 displayer = _metadata.MetadataDisplayer.with_model_file('model.tflite') # 获取标签文件并解析 label_file = displayer.get_packed_associated_file_list()[0] labels = displayer.get_associated_file_buffer(label_file).decode('utf-8').split('\n') print(labels) # 输出即为你设置的类别名称
- 在Android端使用时,通过TFLite Support库读取标签,就能直接获取指定的类别名称。
4. 已训练模型的标签修改(无需重训)
如果已经用错误命名的文件夹完成训练,可直接修改标签文件:
- 导出模型时会生成对应的
labels.txt文件,打开后将内容替换为你需要的类别名称,每行一个。 - 用TFLite Metadata工具将修改后的标签文件重新关联到模型,确保Android端能读取正确标签。
内容的提问来源于stack exchange,提问作者krank s
相关产品推荐
相关产品推荐

