【问题标题】:How to continue training with checkpoints using object_detector.EfficientDetLite4Spec tensorflow lite如何使用 object_detector.EfficientDetLite4Spec tensorflow lite 继续使用检查点进行训练
【发布时间】:2021-11-25 09:41:26
【问题描述】:

我已经在 config.yaml 中设置了我的 EfficientDetLite4 模型 "grad_checkpoint=true"。它已经成功地生成了一些检查点。但是,当我想根据它们继续训练时,我不知道如何使用这些检查点。

每次我训练模型时,它都是从头开始,而不是从我的检查点开始。

下图是我的colab文件系统结构:

<img src="https://i.stack.imgur.com/8EhPx.jpg"/>

my colab file system structure

下图显示了我的检查点存储位置:

<img src="https://i.stack.imgur.com/Ve5al.jpg"/>

model file system here

以下代码显示了我如何配置模型以及如何使用模型进行训练。

import numpy as np
import os

from tflite_model_maker.config import ExportFormat
from tflite_model_maker import model_spec
from tflite_model_maker import object_detector

import tensorflow as tf
assert tf.__version__.startswith('2')

tf.get_logger().setLevel('ERROR')
from absl import logging
logging.set_verbosity(logging.ERROR)

train_data, validation_data, test_data = 
    object_detector.DataLoader.from_csv('csv_path')

spec = object_detector.EfficientDetLite4Spec(
    uri='/content/model',
    model_dir='/content/drive/MyDrive/MathSymbolRecognition/CheckPoints/',
    hparams='grad_checkpoint=true,strategy=gpus',
    epochs=50, batch_size=3,
    steps_per_execution=1, moving_average_decay=0,
    var_freeze_expr='(efficientnet|fpn_cells|resample_p6)',
    tflite_max_detections=25, strategy=spec_strategy
)

model = object_detector.create(train_data, model_spec=spec, batch_size=3, 
    train_whole_model=True, validation_data=validation_data)

【问题讨论】:

    标签: python tensorflow machine-learning tensorflow-lite


    【解决方案1】:

    源代码就是答案!

    我遇到了同样的问题,发现我们传递给 TFLite 模型 Maker 的对象检测器 API 的 model_dir仅用于保存模型的权重:这就是 API 从不从检查点恢复的原因。

    查看此 API 的源代码,我注意到它在内部使用标准的 model.compilemodel.fit 函数,并通过 model.fitcallbacks 参数保存模型的权重。
    这意味着,只要我们能够获得内部 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 重新开始(使用损失值检查恢复是否成功)。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2020-05-22
      • 2020-12-13
      • 2019-04-05
      • 2020-06-17
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多