【发布时间】:2020-07-01 13:31:32
【问题描述】:
ktrain 是深度学习库 TensorFlow Keras(和其他库)的轻量级包装器,可帮助构建、训练和部署神经网络和其他机器学习模型。我可以使用 ktrain 库从检查点恢复训练吗?
【问题讨论】:
ktrain 是深度学习库 TensorFlow Keras(和其他库)的轻量级包装器,可帮助构建、训练和部署神经网络和其他机器学习模型。我可以使用 ktrain 库从检查点恢复训练吗?
【问题讨论】:
是的,你可以。这在 ktrain 常见问题解答中得到了解答。我将在这里复制答案:
# save model and Preprocessor instance after partially training
ktrain.get_predictor(model, preproc).save('/tmp/my_predictor')
# reload Predictor and extract model
model = ktrain.load_predictor('/tmp/my_predictor').model
# re-instantiate Learner and continue training
learner = ktrain.get_learner(model, train_data=trn, val_data=val)
learner.fit_onecycle(2e-5, 1)
注意preproc 这里是一个预处理器 实例。如果使用像texts_from_csv 或images_from_folder 这样的数据加载函数,它将是函数的第三个返回值。或者,如果使用Transformer API进行文本分类,它将是调用text.Transformer的输出(即preproc = text.Transformer('bert-base-uncased', ...))。
transformers 库(如果训练 Hugging Face Transformers 模型)如果模型是 Hugging Face 变形金刚模型,可以直接使用transformers:
# save model using transformers API after partially training
learner.model.save_pretrained('/tmp/my_model')
# reload the model using transformers directly
from transformers import *
model = TFAutoModelForSequenceClassification.from_pretrained('/tmp/my_model')
model.compile(loss='categorical_crossentropy',optimizer='adam', metrics=['accuracy'])
# re-instantiate Learner and continue training
learner = ktrain.get_learner(model, train_data=trn, val_data=val)
learner.fit_onecycle(2e-5, 1)
checkpoint_folder参数保存模型权重checkpoint_folder 参数(例如,learner.autofit(1e-4, 4, checkpoint_folder='/tmp/saved_weights'))仅在每个 epoch 后保存模型的权重。
任何时期的权重都可以使用model.load_weights 方法重新加载到模型中,就像在tf.Keras 中一样。你只需要先重新创建
先说模型。例如,如果训练一个 NER 模型,它的工作原理如下:
# recreate model from scratch
import ktrain
from ktrain import text
model = text.sequence_tagger(...
# load checkpoint weights from 3rd epoch into model
model.load_weights('../models/checkpoints/weights-03.hdf5')
# recreate learner
learner = ktrain.get_learner(model, ...
# continue training here
最后,还有一个 learner.save_model 和 learner.load_model 方法用于在单个会话期间进行交互式训练时保存和重新加载模型。
【讨论】:
learner.autofit(1e-4, 4, checkpoint_folder='/tmp/saved_weights_02')