diff --git a/mediapipe/tasks/python/metadata/metadata_writers/BUILD b/mediapipe/tasks/python/metadata/metadata_writers/BUILD index 268dec49..e86254f2 100644 --- a/mediapipe/tasks/python/metadata/metadata_writers/BUILD +++ b/mediapipe/tasks/python/metadata/metadata_writers/BUILD @@ -57,6 +57,7 @@ py_library( deps = [ ":metadata_info", ":metadata_writer", + ":writer_utils", "//mediapipe/tasks/metadata:image_segmenter_metadata_schema_py", "//mediapipe/tasks/metadata:metadata_schema_py", "//mediapipe/tasks/python/metadata", diff --git a/mediapipe/tasks/python/metadata/metadata_writers/image_segmenter.py b/mediapipe/tasks/python/metadata/metadata_writers/image_segmenter.py index 8e38731a..ee2dc092 100644 --- a/mediapipe/tasks/python/metadata/metadata_writers/image_segmenter.py +++ b/mediapipe/tasks/python/metadata/metadata_writers/image_segmenter.py @@ -22,6 +22,7 @@ from mediapipe.tasks.metadata import metadata_schema_py_generated as _metadata_f from mediapipe.tasks.python.metadata import metadata from mediapipe.tasks.python.metadata.metadata_writers import metadata_info from mediapipe.tasks.python.metadata.metadata_writers import metadata_writer +from mediapipe.tasks.python.metadata.metadata_writers import writer_utils _MODEL_NAME = "ImageSegmenter" @@ -148,10 +149,17 @@ class MetadataWriter(metadata_writer.MetadataWriterBase): writer = metadata_writer.MetadataWriter(model_buffer) writer.add_general_info(_MODEL_NAME, _MODEL_DESCRIPTION) writer.add_image_input(input_norm_mean, input_norm_std) - writer.add_segmentation_output(labels=labels) if activation is not None: option_md = ImageSegmenterOptionsMd(activation) writer.add_custom_metadata(option_md) + num_output_tensors = writer_utils.get_subgraph(model_buffer).OutputsLength() + if num_output_tensors == 2: + # For image segmenter model with 2 output tensors, the first one is + # quality score, and the second one is matting mask. + writer.add_feature_output( + "quality score", "The quality score of matting result." + ) + writer.add_segmentation_output(labels=labels) return cls(writer) def populate(self) -> tuple[bytearray, str]: