|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| """Tests for basnet network."""
|
|
|
| from absl.testing import parameterized
|
| import numpy as np
|
| import tensorflow as tf, tf_keras
|
|
|
| from official.projects.basnet.modeling import basnet_model
|
| from official.projects.basnet.modeling import refunet
|
|
|
|
|
| class BASNetNetworkTest(parameterized.TestCase, tf.test.TestCase):
|
|
|
| @parameterized.parameters(
|
| (256),
|
| (512),
|
| )
|
| def test_basnet_network_creation(
|
| self, input_size):
|
| """Test for creation of a segmentation network."""
|
| inputs = np.random.rand(2, input_size, input_size, 3)
|
| tf_keras.backend.set_image_data_format('channels_last')
|
|
|
| backbone = basnet_model.BASNetEncoder()
|
| decoder = basnet_model.BASNetDecoder()
|
| refinement = refunet.RefUnet()
|
|
|
| model = basnet_model.BASNetModel(
|
| backbone=backbone,
|
| decoder=decoder,
|
| refinement=refinement
|
| )
|
|
|
| sigmoids = model(inputs)
|
| levels = sorted(sigmoids.keys())
|
| self.assertAllEqual(
|
| [2, input_size, input_size, 1],
|
| sigmoids[levels[-1]].numpy().shape)
|
|
|
| def test_serialize_deserialize(self):
|
| """Validate the network can be serialized and deserialized."""
|
| backbone = basnet_model.BASNetEncoder()
|
| decoder = basnet_model.BASNetDecoder()
|
| refinement = refunet.RefUnet()
|
|
|
| model = basnet_model.BASNetModel(
|
| backbone=backbone,
|
| decoder=decoder,
|
| refinement=refinement
|
| )
|
|
|
| config = model.get_config()
|
| new_model = basnet_model.BASNetModel.from_config(config)
|
|
|
|
|
| _ = new_model.to_json()
|
|
|
|
|
| self.assertAllEqual(model.get_config(), new_model.get_config())
|
|
|
|
|
| if __name__ == '__main__':
|
| tf.test.main()
|
|
|