源代码就是答案!
我遇到了同样的问题,发现我们传递给 TFLite 模型 Maker 的对象检测器 API 的 model_dir仅用于保存模型的权重:这就是 API 从不从检查点恢复的原因。
查看此 API 的源代码,我注意到它在内部使用标准的 model.compile 和 model.fit 函数,并通过 model.fit 的 callbacks 参数保存模型的权重。
这意味着,只要我们能够获得内部 keras 模型,我们就可以使用 model.load_weights 来恢复我们的检查点!
如果您想了解更多关于我在下面使用的一些功能的作用,这些是指向源代码的链接:
这是代码:
#Old useful imports
import tensorflow as tf
from tflite_model_maker.config import QuantizationConfig
from tflite_model_maker.config import ExportFormat
from tflite_model_maker import model_spec
from tflite_model_maker import object_detector
from tflite_model_maker.object_detector import DataLoader
#Import the same libs that TFLiteModelMaker interally uses
from tensorflow_examples.lite.model_maker.third_party.efficientdet.keras import train
from tensorflow_examples.lite.model_maker.third_party.efficientdet.keras import train_lib
#Setup variables
batch_size = 1
epochs = 50
checkpoint_dir = "/content/..." #whataver you checkpoint dir is
#Create whichever object detector's spec you want
spec = object_detector.EfficientDetLite4Spec(
model_name='efficientdet-lite4',
uri='https://tfhub.dev/tensorflow/efficientdet/lite4/feature-vector/2',
hparams='', #enable grad_checkpoint=True if you want
model_dir=checkpoint_dir,
epochs=epochs,
batch_size=batch_size,
steps_per_execution=1,
moving_average_decay=0,
var_freeze_expr='(efficientnet|fpn_cells|resample_p6)',
tflite_max_detections=25,
strategy=None,
tpu=None,
gcp_project=None,
tpu_zone=None,
use_xla=False,
profile=False,
debug=False,
tf_random_seed=111111,
verbose=1
)
#Load you datasets as usual
train_data, validation_data, test_data = object_detector.DataLoader.from_csv('/path/to/csv.csv')
#Create the object detector as described in the documentation
detector = object_detector.create(train_data,
model_spec=spec,
batch_size=batch_size,
train_whole_model=True,
validation_data=validation_data,
epochs = epochs,
do_train = False
)
"""
From here on we use internal/"private" functions of the API,
you can tell because the methods's names begin with an underscore
"""
#Convert the datasets for training
train_ds, steps_per_epoch, _ = detector._get_dataset_and_steps(train_data, batch_size, is_training = True)
validation_ds, validation_steps, val_json_file = detector._get_dataset_and_steps(validation_data, batch_size, is_training = False)
#Get the interal keras model
model = detector.create_model()
#Copy what the API interally does as setup
config = spec.config
config.update(
dict(
steps_per_epoch=steps_per_epoch,
eval_samples=batch_size * validation_steps,
val_json_file=val_json_file,
batch_size=batch_size
)
)
train.setup_model(model, config) #This is the model.compile call basically
model.summary()
#Load the weights from the latest checkpoint
try:
latest = tf.train.latest_checkpoint(checkpoint_dir)
model.load_weights(latest)
print("Checkpoint found {}".format(latest))
except Exception as e:
print("Checkpoint not found: ", e)
#Train the model
model.fit(
train_ds,
epochs=epochs,
steps_per_epoch=steps_per_epoch,
validation_data=validation_ds,
validation_steps=validation_steps,
callbacks=train_lib.get_callbacks(config.as_dict(), validation_ds) #This is for saving checkpoints at the end of every epoch
)
#Save/export the trained model
export_dir = "/content/..." #save the tflite wherever you want
quant_config = QuantizationConfig.for_float16() #or whatever quantization you want
detector.model = model #inject our trained model into the object detector
detector.export(export_dir = export_dir, tflite_filename='model.tflite', quantization_config = quant_config)
请注意,检查点不保存纪元的信息。
这意味着即使您的权重将恢复,您的历元计数器也将始终从 1 重新开始(使用损失值检查恢复是否成功)。