|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| """Export module for BASNet."""
|
|
|
| import tensorflow as tf, tf_keras
|
|
|
| from official.projects.basnet.tasks import basnet
|
| from official.vision.serving import semantic_segmentation
|
|
|
|
|
| class BASNetModule(semantic_segmentation.SegmentationModule):
|
| """BASNet Module."""
|
|
|
| def _build_model(self):
|
| input_specs = tf_keras.layers.InputSpec(
|
| shape=[self._batch_size] + self._input_image_size + [3])
|
|
|
| return basnet.build_basnet_model(
|
| input_specs=input_specs,
|
| model_config=self.params.task.model,
|
| l2_regularizer=None)
|
|
|
| def serve(self, images):
|
| """Cast image to float and run inference.
|
|
|
| Args:
|
| images: uint8 Tensor of shape [batch_size, None, None, 3]
|
| Returns:
|
| Tensor holding classification output logits.
|
| """
|
| with tf.device('cpu:0'):
|
| images = tf.cast(images, dtype=tf.float32)
|
|
|
| images = tf.nest.map_structure(
|
| tf.identity,
|
| tf.map_fn(
|
| self._build_inputs, elems=images,
|
| fn_output_signature=tf.TensorSpec(
|
| shape=self._input_image_size + [3], dtype=tf.float32),
|
| parallel_iterations=32
|
| )
|
| )
|
|
|
| masks = self.inference_step(images)
|
| keys = sorted(masks.keys())
|
| output = tf.image.resize(
|
| masks[keys[-1]],
|
| self._input_image_size, method='bilinear')
|
|
|
| return dict(predicted_masks=output)
|
|
|