Add a test for custom load weights in BERT
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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"])
|
||||
|
||||
Reference in New Issue
Block a user