|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| """Tests for official.nlp.projects.bigbird.encoder."""
|
|
|
| import numpy as np
|
| import tensorflow as tf, tf_keras
|
|
|
| from official.projects.bigbird import encoder
|
|
|
|
|
| class BigBirdEncoderTest(tf.test.TestCase):
|
|
|
| def test_encoder(self):
|
| sequence_length = 1024
|
| batch_size = 2
|
| vocab_size = 1024
|
| network = encoder.BigBirdEncoder(
|
| num_layers=1, vocab_size=1024, max_position_embeddings=4096)
|
| word_id_data = np.random.randint(
|
| vocab_size, size=(batch_size, sequence_length))
|
| mask_data = np.random.randint(2, size=(batch_size, sequence_length))
|
| type_id_data = np.random.randint(2, size=(batch_size, sequence_length))
|
| outputs = network([word_id_data, mask_data, type_id_data])
|
| self.assertEqual(outputs["sequence_output"].shape,
|
| (batch_size, sequence_length, 768))
|
|
|
| def test_save_restore(self):
|
| sequence_length = 1024
|
| batch_size = 2
|
| vocab_size = 1024
|
| network = encoder.BigBirdEncoder(
|
| num_layers=1, vocab_size=1024, max_position_embeddings=4096)
|
| word_id_data = np.random.randint(
|
| vocab_size, size=(batch_size, sequence_length))
|
| mask_data = np.random.randint(2, size=(batch_size, sequence_length))
|
| type_id_data = np.random.randint(2, size=(batch_size, sequence_length))
|
| inputs = dict(
|
| input_word_ids=word_id_data,
|
| input_mask=mask_data,
|
| input_type_ids=type_id_data)
|
| ref_outputs = network(inputs)
|
| model_path = self.get_temp_dir() + "/model"
|
| network.save(model_path)
|
| loaded = tf_keras.models.load_model(model_path)
|
| outputs = loaded(inputs)
|
| self.assertAllClose(outputs["sequence_output"],
|
| ref_outputs["sequence_output"])
|
|
|
|
|
| if __name__ == "__main__":
|
| tf.test.main()
|
|
|