import os
import json
import tensorflow as tf
from Object_Category_Number import ObjCategory
assert tf.__version__.startswith('2')
from mediapipe_model_maker import object_detector
from mediapipe_model_maker import quantization
train_dataset_name = "train_images/mask"
train_dataset_path = train_dataset_name + "/train"
validation_dataset_path = train_dataset_name + "/validation"
with open(os.path.join(train_dataset_path, "labels.json"), "r") as f:
labels_json = json.load(f)
categoryList = []
for category_item in labels_json["categories"]:
categoryList.append(ObjCategory(category_item['name'], category_item['id'], 0))
for annotation_item in labels_json["annotations"]:
for cc in categoryList:
if cc.categoryId == annotation_item['category_id']:
cc.categoryNum = cc.categoryNum + 1
for cc in categoryList:
print(f"{cc.categoryId}: {cc.categoryName} : {cc.categoryNum} ")
train_data = object_detector.Dataset.from_coco_folder(train_dataset_path, cache_dir="/tmp/od_data/train")
validation_data = object_detector.Dataset.from_coco_folder(validation_dataset_path, cache_dir="/tmp/od_data/validation")
print("train_data size: ", train_data.size)
print("validation_data size: ", validation_data.size)
spec = object_detector.SupportedModels.MOBILENET_MULTI_AVG_I384
hparamsTest = object_detector.HParams(export_dir=train_dataset_name)
hparams = object_detector.HParams(export_dir=train_dataset_name, learning_rate=0.3, batch_size=8, epochs=50)
options = object_detector.ObjectDetectorOptions(
supported_model=spec,
hparams=hparams
)
model = object_detector.ObjectDetector.create(
train_data=train_data,
validation_data=validation_data,
options=options)
loss, coco_metrics = model.evaluate(validation_data, batch_size=8)
print(f"Validation loss: {loss}")
print(f"Validation coco metrics: {coco_metrics}")
lastDirName = os.path.basename(train_dataset_name)
model.export_model(lastDirName + '.tflite')
print(f"---------- model.export_model 已经导出模型 ,准备开始Model quantization 工作了 ------------")
quantization_config = quantization.QuantizationConfig.for_float16()
model.restore_float_ckpt()
model.export_model(model_name=lastDirName + "_fp16.tflite", quantization_config=quantization_config)
print(f"----------- model.export_model 已经导出模型 XXX_fp16.tflite ---------------")