Project import generated by Copybara.
GitOrigin-RevId: 73d686c40057684f8bfaca285368bf1813f9fc26
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user