Add a test for custom load weights in BERT

This commit is contained in:
Julien Plu
2020-10-05 14:14:35 +02:00
parent f40dbd795b
commit 03117e929c
2 changed files with 9 additions and 6 deletions
+1 -1
View File
@@ -750,7 +750,7 @@ class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFGenerationMixin):
)
if output_loading_info:
loading_info = {"missing_keys": missing_keys, "unexpected_layers_weights": unexpected_keys}
loading_info = {"missing_keys": missing_keys, "unexpected_keys": unexpected_keys}
return model, loading_info
+8 -5
View File
@@ -317,9 +317,12 @@ class TFBertModelTest(TFModelTesterMixin, unittest.TestCase):
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_bert_for_token_classification(*config_and_inputs)
@slow
def test_model_from_pretrained(self):
# for model_name in TF_BERT_PRETRAINED_MODEL_ARCHIVE_LIST[:1]:
for model_name in ["bert-base-uncased"]:
model = TFBertModel.from_pretrained(model_name)
self.assertIsNotNone(model)
model = TFBertModel.from_pretrained("jplu/tiny-tf-bert-random")
self.assertIsNotNone(model)
def test_custom_load_tf_weights(self):
model, output_loading_info = TFBertForTokenClassification.from_pretrained("jplu/tiny-tf-bert-random", use_cdn=False, output_loading_info=True)
self.assertEqual(sorted(output_loading_info["unexpected_keys"]), ['mlm___cls', 'nsp___cls'])
for layer in output_loading_info["missing_keys"]:
self.assertTrue(layer.split("_")[0] in ["dropout", "classifier"])