Project import generated by Copybara.

GitOrigin-RevId: 73d686c40057684f8bfaca285368bf1813f9fc26
This commit is contained in:
MediaPipe Team
2022-03-21 12:12:39 -07:00
committed by jqtang
parent e6c19885c6
commit cc6a2f7af6
266 changed files with 3658 additions and 1681 deletions
+1 -1
View File
@@ -59,7 +59,7 @@ cc_library(
"//mediapipe/calculators/core:gate_calculator",
"//mediapipe/calculators/core:pass_through_calculator",
"//mediapipe/calculators/core:side_packet_to_stream_calculator",
"//mediapipe/calculators/core:split_landmarks_calculator",
"//mediapipe/calculators/core:split_proto_list_calculator",
"//mediapipe/calculators/core:string_to_int_calculator",
"//mediapipe/calculators/image:image_transformation_calculator",
"//mediapipe/calculators/util:detection_unique_id_calculator",
+1 -1
View File
@@ -392,7 +392,7 @@ void CalculatorGraphSubmodule(pybind11::module* module) {
}
return std::string();
},
R"doc(Combines error messages as a single std::string.
R"doc(Combines error messages as a single string.
Examples:
if graph.has_error():
+8 -8
View File
@@ -80,13 +80,13 @@ void PublicPacketCreators(pybind11::module* m) {
m->def(
"create_string",
[](const std::string& data) { return MakePacket<std::string>(data); },
R"doc(Create a MediaPipe std::string Packet from a str.
R"doc(Create a MediaPipe string Packet from a str.
Args:
data: A str.
Returns:
A MediaPipe std::string Packet.
A MediaPipe string Packet.
Raises:
TypeError: If the input is not a str.
@@ -100,13 +100,13 @@ void PublicPacketCreators(pybind11::module* m) {
m->def(
"create_string",
[](const py::bytes& data) { return MakePacket<std::string>(data); },
R"doc(Create a MediaPipe std::string Packet from a bytes object.
R"doc(Create a MediaPipe string Packet from a bytes object.
Args:
data: A bytes object.
Returns:
A MediaPipe std::string Packet.
A MediaPipe string Packet.
Raises:
TypeError: If the input is not a bytes object.
@@ -498,13 +498,13 @@ void PublicPacketCreators(pybind11::module* m) {
[](const std::vector<std::string>& data) {
return MakePacket<std::vector<std::string>>(data);
},
R"doc(Create a MediaPipe std::string vector Packet from a list of str.
R"doc(Create a MediaPipe string vector Packet from a list of str.
Args:
data: A list of str.
Returns:
A MediaPipe std::string vector Packet.
A MediaPipe string vector Packet.
Raises:
TypeError: If the input is not a list of str.
@@ -546,7 +546,7 @@ void PublicPacketCreators(pybind11::module* m) {
[](const std::map<std::string, Packet>& data) {
return MakePacket<std::map<std::string, Packet>>(data);
},
R"doc(Create a MediaPipe std::string to packet map Packet from a dictionary.
R"doc(Create a MediaPipe string to packet map Packet from a dictionary.
Args:
data: A dictionary that has (str, Packet) pairs.
@@ -561,7 +561,7 @@ void PublicPacketCreators(pybind11::module* m) {
dict_packet = mp.packet_creator.create_string_to_packet_map({
'float': mp.packet_creator.create_float(0.1),
'int': mp.packet_creator.create_int(1),
'std::string': mp.packet_creator.create_string('1')
'string': mp.packet_creator.create_string('1')
data = mp.packet_getter.get_str_to_packet_dict(dict_packet)
)doc",
py::arg().noconvert(), py::return_value_policy::move);
+13 -13
View File
@@ -42,16 +42,16 @@ namespace py = pybind11;
void PublicPacketGetters(pybind11::module* m) {
m->def("get_str", &GetContent<std::string>,
R"doc(Get the content of a MediaPipe std::string Packet as a str.
R"doc(Get the content of a MediaPipe string Packet as a str.
Args:
packet: A MediaPipe std::string Packet.
packet: A MediaPipe string Packet.
Returns:
A str.
Raises:
ValueError: If the Packet doesn't contain std::string data.
ValueError: If the Packet doesn't contain string data.
Examples:
packet = mp.packet_creator.create_string('abc')
@@ -63,16 +63,16 @@ void PublicPacketGetters(pybind11::module* m) {
[](const Packet& packet) {
return py::bytes(GetContent<std::string>(packet));
},
R"doc(Get the content of a MediaPipe std::string Packet as a bytes object.
R"doc(Get the content of a MediaPipe string Packet as a bytes object.
Args:
packet: A MediaPipe std::string Packet.
packet: A MediaPipe string Packet.
Returns:
A bytes object.
Raises:
ValueError: If the Packet doesn't contain std::string data.
ValueError: If the Packet doesn't contain string data.
Examples:
packet = mp.packet_creator.create_string(b'\xd0\xd0\xd0')
@@ -266,7 +266,7 @@ void PublicPacketGetters(pybind11::module* m) {
m->def(
"get_str_list", &GetContent<std::vector<std::string>>,
R"doc(Get the content of a MediaPipe std::string vector Packet as a str list.
R"doc(Get the content of a MediaPipe string vector Packet as a str list.
Args:
packet: A MediaPipe Packet that holds std:vector<std::string>.
@@ -322,7 +322,7 @@ void PublicPacketGetters(pybind11::module* m) {
dict_packet = mp.packet_creator.create_string_to_packet_map({
'float': packet_creator.create_float(0.1),
'int': packet_creator.create_int(1),
'std::string': packet_creator.create_string('1')
'string': packet_creator.create_string('1')
data = mp.packet_getter.get_str_to_packet_dict(dict_packet)
)doc");
@@ -418,11 +418,11 @@ void InternalPacketGetters(pybind11::module* m) {
m->def(
"_get_serialized_proto",
[](const Packet& packet) {
// By default, py::bytes is an extra copy of the original std::string
// object: https://github.com/pybind/pybind11/issues/1236 However, when
// Pybind11 performs the C++ to Python transition, it only increases the
// py::bytes object's ref count. See the implmentation at line 1583 in
// "pybind11/cast.h".
// By default, py::bytes is an extra copy of the original string object:
// https://github.com/pybind/pybind11/issues/1236
// However, when Pybind11 performs the C++ to Python transition, it
// only increases the py::bytes object's ref count. See the
// implmentation at line 1583 in "pybind11/cast.h".
return py::bytes(packet.GetProtoMessageLite().SerializeAsString());
},
py::return_value_policy::move);
+80 -60
View File
@@ -79,9 +79,16 @@ CALCULATOR_TO_OPTIONS = {
}
def type_names_from_oneof(oneof_type_name: str) -> Optional[List[str]]:
if oneof_type_name.startswith('OneOf<') and oneof_type_name.endswith('>'):
comma_separated_types = oneof_type_name[len('OneOf<'):-len('>')]
return [n.strip() for n in comma_separated_types.split(',')]
return None
# TODO: Support more packet data types, such as "Any" type.
@enum.unique
class _PacketDataType(enum.Enum):
class PacketDataType(enum.Enum):
"""The packet data types supported by the SolutionBase class."""
STRING = 'string'
BOOL = 'bool'
@@ -96,79 +103,86 @@ class _PacketDataType(enum.Enum):
PROTO_LIST = 'proto_list'
@staticmethod
def from_registered_name(registered_name: str) -> '_PacketDataType':
return NAME_TO_TYPE[registered_name]
def from_registered_name(registered_name: str) -> 'PacketDataType':
try:
return NAME_TO_TYPE[registered_name]
except KeyError as e:
names = type_names_from_oneof(registered_name)
if names:
for n in names:
if n in NAME_TO_TYPE.keys():
return NAME_TO_TYPE[n]
raise e
NAME_TO_TYPE: Mapping[str, '_PacketDataType'] = {
NAME_TO_TYPE: Mapping[str, 'PacketDataType'] = {
'string':
_PacketDataType.STRING,
PacketDataType.STRING,
'bool':
_PacketDataType.BOOL,
PacketDataType.BOOL,
'::std::vector<bool>':
_PacketDataType.BOOL_LIST,
PacketDataType.BOOL_LIST,
'int':
_PacketDataType.INT,
PacketDataType.INT,
'float':
_PacketDataType.FLOAT,
PacketDataType.FLOAT,
'::std::vector<float>':
_PacketDataType.FLOAT_LIST,
PacketDataType.FLOAT_LIST,
'::mediapipe::Matrix':
_PacketDataType.AUDIO,
PacketDataType.AUDIO,
'::mediapipe::ImageFrame':
_PacketDataType.IMAGE_FRAME,
PacketDataType.IMAGE_FRAME,
'::mediapipe::Classification':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::ClassificationList':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::ClassificationListCollection':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::Detection':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::DetectionList':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::Landmark':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::LandmarkList':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::LandmarkListCollection':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::NormalizedLandmark':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::FrameAnnotation':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::Trigger':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::Rect':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::NormalizedRect':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::NormalizedLandmarkList':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::NormalizedLandmarkListCollection':
_PacketDataType.PROTO,
PacketDataType.PROTO,
'::mediapipe::Image':
_PacketDataType.IMAGE,
PacketDataType.IMAGE,
'::std::vector<::mediapipe::Classification>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::ClassificationList>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::Detection>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::DetectionList>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::Landmark>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::LandmarkList>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::NormalizedLandmark>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::NormalizedLandmarkList>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::Rect>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
'::std::vector<::mediapipe::NormalizedRect>':
_PacketDataType.PROTO_LIST,
PacketDataType.PROTO_LIST,
}
@@ -196,7 +210,8 @@ class SolutionBase:
graph_config: Optional[calculator_pb2.CalculatorGraphConfig] = None,
calculator_params: Optional[Mapping[str, Any]] = None,
side_inputs: Optional[Mapping[str, Any]] = None,
outputs: Optional[List[str]] = None):
outputs: Optional[List[str]] = None,
stream_type_hints: Optional[Mapping[str, PacketDataType]] = None):
"""Initializes the SolutionBase object.
Args:
@@ -209,6 +224,7 @@ class SolutionBase:
outputs: A list of the graph output stream names to observe. If the list
is empty, all the output streams listed in the graph config will be
automatically observed by default.
stream_type_hints: A mapping from the stream name to its packet type hint.
Raises:
FileNotFoundError: If the binary graph file can't be found.
@@ -240,7 +256,7 @@ class SolutionBase:
validated_graph.initialize(graph_config=graph_config)
canonical_graph_config_proto = self._initialize_graph_interface(
validated_graph, side_inputs, outputs)
validated_graph, side_inputs, outputs, stream_type_hints)
if calculator_params:
self._modify_calculator_options(canonical_graph_config_proto,
calculator_params)
@@ -310,15 +326,15 @@ class SolutionBase:
self._simulated_timestamp += 33333
for stream_name, data in input_dict.items():
input_stream_type = self._input_stream_type_info[stream_name]
if (input_stream_type == _PacketDataType.PROTO_LIST or
input_stream_type == _PacketDataType.AUDIO):
if (input_stream_type == PacketDataType.PROTO_LIST or
input_stream_type == PacketDataType.AUDIO):
# TODO: Support audio data.
raise NotImplementedError(
f'SolutionBase can only process non-audio and non-proto-list data. '
f'{self._input_stream_type_info[stream_name].name} '
f'type is not supported yet.')
elif (input_stream_type == _PacketDataType.IMAGE_FRAME or
input_stream_type == _PacketDataType.IMAGE):
elif (input_stream_type == PacketDataType.IMAGE_FRAME or
input_stream_type == PacketDataType.IMAGE):
if data.shape[2] != RGB_CHANNELS:
raise ValueError('Input image must contain three channel rgb data.')
self._graph.add_packet_to_input_stream(
@@ -364,7 +380,8 @@ class SolutionBase:
self,
validated_graph: validated_graph_config.ValidatedGraphConfig,
side_inputs: Optional[Mapping[str, Any]] = None,
outputs: Optional[List[str]] = None):
outputs: Optional[List[str]] = None,
stream_type_hints: Optional[Mapping[str, PacketDataType]] = None):
"""Gets graph interface type information and returns the canonical graph config proto."""
canonical_graph_config_proto = calculator_pb2.CalculatorGraphConfig()
@@ -375,13 +392,16 @@ class SolutionBase:
return tag_index_name.split(':')[-1]
# Gets the packet type information of the input streams and output streams
# from the validated calculator graph. The mappings from the stream names to
# the packet data types is for deciding which packet creator and getter
# methods to call in the process() method.
# from the user provided stream_type_hints field or validated calculator
# graph. The mappings from the stream names to the packet data types is
# for deciding which packet creator and getter methods to call in the
# process() method.
def get_stream_packet_type(packet_tag_index_name):
return _PacketDataType.from_registered_name(
validated_graph.registered_stream_type_name(
get_name(packet_tag_index_name)))
stream_name = get_name(packet_tag_index_name)
if stream_type_hints and stream_name in stream_type_hints.keys():
return stream_type_hints[stream_name]
return PacketDataType.from_registered_name(
validated_graph.registered_stream_type_name(stream_name))
self._input_stream_type_info = {
get_name(tag_index_name): get_stream_packet_type(tag_index_name)
@@ -402,7 +422,7 @@ class SolutionBase:
# packet data types is for making the input_side_packets dict for graph
# start_run().
def get_side_packet_type(packet_tag_index_name):
return _PacketDataType.from_registered_name(
return PacketDataType.from_registered_name(
validated_graph.registered_side_packet_type_name(
get_name(packet_tag_index_name)))
@@ -503,16 +523,16 @@ class SolutionBase:
if num_modified < len(nested_calculator_params):
raise ValueError('Not all calculator params are valid.')
def _make_packet(self, packet_data_type: _PacketDataType,
def _make_packet(self, packet_data_type: PacketDataType,
data: Any) -> packet.Packet:
if (packet_data_type == _PacketDataType.IMAGE_FRAME or
packet_data_type == _PacketDataType.IMAGE):
if (packet_data_type == PacketDataType.IMAGE_FRAME or
packet_data_type == PacketDataType.IMAGE):
return getattr(packet_creator, 'create_' + packet_data_type.value)(
data, image_format=image_frame.ImageFormat.SRGB)
else:
return getattr(packet_creator, 'create_' + packet_data_type.value)(data)
def _get_packet_content(self, packet_data_type: _PacketDataType,
def _get_packet_content(self, packet_data_type: PacketDataType,
output_packet: packet.Packet) -> Any:
"""Gets packet content from a packet by type.
@@ -527,10 +547,10 @@ class SolutionBase:
if output_packet.is_empty():
return None
if packet_data_type == _PacketDataType.STRING:
if packet_data_type == PacketDataType.STRING:
return packet_getter.get_str(output_packet)
elif (packet_data_type == _PacketDataType.IMAGE_FRAME or
packet_data_type == _PacketDataType.IMAGE):
elif (packet_data_type == PacketDataType.IMAGE_FRAME or
packet_data_type == PacketDataType.IMAGE):
return getattr(packet_getter, 'get_' +
packet_data_type.value)(output_packet).numpy_view()
else:
+29
View File
@@ -22,6 +22,7 @@ from google.protobuf import text_format
from mediapipe.framework import calculator_pb2
from mediapipe.framework.formats import detection_pb2
from mediapipe.python import solution_base
from mediapipe.python.solution_base import PacketDataType
CALCULATOR_OPTIONS_TEST_GRAPH_CONFIG = """
input_stream: 'image_in'
@@ -348,6 +349,34 @@ class SolutionBaseTest(parameterized.TestCase):
self.assertTrue(np.array_equal(input_image, outputs.image_out))
solution.reset()
def test_solution_stream_type_hints(self):
text_config = """
input_stream: 'union_type_image_in'
output_stream: 'image_type_out'
node {
calculator: 'ToImageCalculator'
input_stream: 'IMAGE:union_type_image_in'
output_stream: 'IMAGE:image_type_out'
}
"""
config_proto = text_format.Parse(text_config,
calculator_pb2.CalculatorGraphConfig())
input_image = np.arange(27, dtype=np.uint8).reshape(3, 3, 3)
with solution_base.SolutionBase(
graph_config=config_proto,
stream_type_hints={'union_type_image_in': PacketDataType.IMAGE
}) as solution:
for _ in range(20):
outputs = solution.process(input_image)
self.assertTrue(np.array_equal(input_image, outputs.image_type_out))
with solution_base.SolutionBase(
graph_config=config_proto,
stream_type_hints={'union_type_image_in': PacketDataType.IMAGE_FRAME
}) as solution2:
for _ in range(20):
outputs = solution2.process(input_image)
self.assertTrue(np.array_equal(input_image, outputs.image_type_out))
def _process_and_verify(self,
config_proto,
side_inputs=None,