diff --git a/docs/Reduced_Operator_Kernel_build.md b/docs/Reduced_Operator_Kernel_build.md index 7c6413284b..02e761e5c6 100644 --- a/docs/Reduced_Operator_Kernel_build.md +++ b/docs/Reduced_Operator_Kernel_build.md @@ -2,13 +2,13 @@ In order to reduce the compiled binary size of ONNX Runtime (ORT), the operator kernels included in the build can be reduced to just the kernels required by your model/s. -A configuration file must be created with details of the kernels that are required. +A configuration file must be created with details of the kernels that are required. -Following that, ORT must be manually built, providing the configuration file in the `--include_ops_by_config` parameter. The build process will update the ORT kernel registration source files to exclude the unused kernels. +Following that, ORT must be manually built, providing the configuration file in the `--include_ops_by_config` parameter. The build process will update the ORT kernel registration source files to exclude the unused kernels. See the [build instructions](https://www.onnxruntime.ai/docs/how-to/build.html#build-instructions) for more details on building ORT. -When building ORT with a reduced set of kernel registrations, `--skip_tests` **MUST** be specified as the kernel reduction will render many of the unit tests invalid. +When building ORT with a reduced set of kernel registrations, `--skip_tests` **MUST** be specified as the kernel reduction will render many of the unit tests invalid. NOTE: The operator exclusion logic when building with an operator reduction configuration file will only disable kernel registrations each time it runs. It will NOT re-enable previously disabled kernels. If you wish to change the list of kernels included, it is best to revert the repository to a clean state (e.g. via `git reset --hard`) before building ORT again. @@ -75,7 +75,7 @@ If, for example, the types of inputs 0 and 1 were important, the entry may look `{"inputs": {"0": ["float", "int32_t"], "1": ["int32_t"]}}` Finally some operators do non-standard things and store their type information under a 'custom' key. -ai.onnx.OneHot is an example of this, where 3 type names from the inputs are combined into a string. - `{"custom": ["float_int64_t_int64_t", "int64_t_string_int64_t"]}` +ai.onnx.OneHot is an example of this, where the three input types are combined into a triple. + `{"custom": [["float", "int64_t", "int64_t"], ["int64_t", "std::string", "int64_t"]]}` -For these reasons, it is best to generate the configuration file first, and manually edit any entries if needed. +For these reasons, it is best to generate the configuration file first, and manually edit any entries if needed. diff --git a/tools/ci_build/github/azure-pipelines/linux-cpu-minimal-build-ci-pipeline.yml b/tools/ci_build/github/azure-pipelines/linux-cpu-minimal-build-ci-pipeline.yml index 7fe6b69152..9b45ad9e04 100644 --- a/tools/ci_build/github/azure-pipelines/linux-cpu-minimal-build-ci-pipeline.yml +++ b/tools/ci_build/github/azure-pipelines/linux-cpu-minimal-build-ci-pipeline.yml @@ -130,7 +130,9 @@ jobs: displayName: Build minimal onnxruntime [exceptions ENABLED, type reduction ENABLED (globally allowed types)] and run tests inputs: script: | - echo "!globally_allowed_types;bool,float,int8_t,uint8_t" \ + printf "%s\n%s\n" \ + "!globally_allowed_types;bool,float,int8_t,uint8_t" \ + "!no_ops_specified_means_all_ops_are_required" \ > $(test_data_directory)/globally_allowed_types.config && \ docker run --rm \ --volume $(Build.SourcesDirectory):/onnxruntime_src \ diff --git a/tools/ci_build/op_registration_utils.py b/tools/ci_build/op_registration_utils.py index 7de949340c..7ea776d00c 100644 --- a/tools/ci_build/op_registration_utils.py +++ b/tools/ci_build/op_registration_utils.py @@ -184,6 +184,11 @@ def _process_lines(lines: typing.List[str], offset: int, registration_processor: registration_processor.process_registration(lines_to_process, domain, op_type, int(start_version), int(end_version), type) + else: + log.warning("Ignoring unhandled kernel registration variant: {}".format(code_line)) + for line in lines_to_process: + registration_processor.process_other_line(line) + return offset + 1 diff --git a/tools/ci_build/reduce_op_kernels.py b/tools/ci_build/reduce_op_kernels.py index 8387a541b8..427bb77f7b 100755 --- a/tools/ci_build/reduce_op_kernels.py +++ b/tools/ci_build/reduce_op_kernels.py @@ -5,7 +5,6 @@ import argparse import op_registration_utils import os -import re import shutil import sys import typing @@ -19,74 +18,24 @@ ort_tools_py_path = os.path.abspath(os.path.join(ort_root, 'tools', 'python')) sys.path.append(ort_tools_py_path) from util import parse_config # noqa -from util.ort_format_model.operator_type_usage_processors import OperatorTypeUsageManager # noqa +from util.ort_format_model.operator_type_usage_processors import OpTypeImplFilterInterface # noqa log = get_logger("reduce_op_kernels") -# valid C++ scalar types that can be specified as globally allowed types -_valid_allowed_types = { - "bool", - "uint8_t", "uint16_t", "uint32_t", "uint64_t", - "int8_t", "int16_t", "int32_t", "int64_t", - "MLFloat16", "BFloat16", # in onnxruntime namespace - "float", "double", - "string", # in std namespace -} - - -def _validated_globally_allowed_types(globally_allowed_types: typing.Collection[str]) -> typing.Set[str]: - '''Return a valid set of globally allowed types.''' - # ensure globally_allowed_types is a set - if not isinstance(globally_allowed_types, set): - globally_allowed_types = set(globally_allowed_types) - - if not globally_allowed_types <= _valid_allowed_types: - raise ValueError( - "Globally allowed types must be a subset of valid allowed types. Actual: {}, valid: {}".format( - globally_allowed_types, sorted(_valid_allowed_types))) - - return globally_allowed_types - - -def _type_re_from_globally_allowed_types(globally_allowed_types: typing.Set[str]) -> typing.re.Pattern: - '''Return a regular expression to match type registration strings to a set of globally allowed types.''' - # to keep a registration, the type should match patterns like: - # 1. T0 - # 2. T0_T1_T2 - # where Ti is a member of globally_allowed_types and multiple Ti's are delimited by "_" - # this covers both the common case (1) and special cases like OneHot registration (2) - allowed_type_subpattern = \ - "(?:" + "|".join(re.escape(allowed_type) for allowed_type in sorted(globally_allowed_types)) + ")" - return re.compile("^{0}(?:_{0})*$".format(allowed_type_subpattern)) - - class _ExcludingRegistrationProcessor(op_registration_utils.RegistrationProcessor): '''Registration processor that excludes registrations and writes the result to an output file.''' - def __init__(self, required_ops: dict, op_type_usage_manager: typing.Optional[OperatorTypeUsageManager], - globally_allowed_types: typing.Optional[typing.Set[str]], output_file: str): + def __init__(self, required_ops: typing.Optional[dict], + op_type_impl_filter: typing.Optional[OpTypeImplFilterInterface], + output_file: str): self._required_ops = required_ops - - if op_type_usage_manager is not None and globally_allowed_types is not None: - raise ValueError("At most one of op_type_usage_manager and globally_allowed_types may be provided.") - - self._op_type_usage_manager = op_type_usage_manager - - self._enable_all_ops = globally_allowed_types is not None and not required_ops - if self._enable_all_ops: - log.info("No required ops were specified but globally allowed types were specified. " - "Globally allowed types will be used to exclude op implementations.") - - self._globally_allowed_types_re = \ - _type_re_from_globally_allowed_types(globally_allowed_types) \ - if globally_allowed_types is not None else None - + self._op_type_impl_filter = op_type_impl_filter self._output_file = output_file def _is_op_required(self, domain: str, operator: str, start_version: int, end_version: typing.Optional[int]) -> typing.Tuple[bool, str]: - '''See if an op should be excluded because it is not required.''' - if self._enable_all_ops: + '''See if an op is required.''' + if self._required_ops is None: return True if domain not in self._required_ops: @@ -116,17 +65,10 @@ class _ExcludingRegistrationProcessor(op_registration_utils.RegistrationProcesso exclude = True reason = "Entire op is not required." - if not exclude and type is not None: - if self._op_type_usage_manager is not None: - if not self._op_type_usage_manager.is_typed_registration_needed(domain, operator, type): - exclude = True - reason = "Specific typed registration is not required." - - elif self._globally_allowed_types_re is not None: - if not self._globally_allowed_types_re.match(type): - exclude = True - reason = "Specific typed registration does not contain globally allowed types." - + if not exclude and type is not None and self._op_type_impl_filter is not None: + if not self._op_type_impl_filter.is_typed_registration_needed(domain, operator, type): + exclude = True + reason = "Specific typed registration is not required." else: log.warning('Keeping {} registration from unknown domain: {}' .format(registration_identifier, constant_for_domain)) @@ -152,9 +94,8 @@ class _ExcludingRegistrationProcessor(op_registration_utils.RegistrationProcesso def _process_provider_registrations( ort_root: str, use_cuda: bool, - required_ops: dict, - op_type_usage_manager: typing.Optional[OperatorTypeUsageManager], - globally_allowed_types: typing.Optional[typing.Set[str]]): + required_ops: typing.Optional[dict], + op_type_impl_filter: typing.Optional[OpTypeImplFilterInterface]): '''Rewrite provider registration files.''' kernel_registration_files = op_registration_utils.get_kernel_registration_files(ort_root, use_cuda) @@ -169,8 +110,7 @@ def _process_provider_registrations( # read from backup and overwrite original with commented out lines for any kernels that are not required with open(kernel_registration_file, 'w') as file_to_write: - processor = _ExcludingRegistrationProcessor( - required_ops, op_type_usage_manager, globally_allowed_types, file_to_write) + processor = _ExcludingRegistrationProcessor(required_ops, op_type_impl_filter, file_to_write) op_registration_utils.process_kernel_registration_file(backup_path, processor) @@ -231,20 +171,11 @@ def reduce_ops(config_path: str, enable_type_reduction: bool = False, use_cuda: :param enable_type_reduction: Whether per operator type reduction is enabled :param use_cuda: Whether to reduce op kernels for the CUDA provider ''' - required_ops, op_type_usage_manager, globally_allowed_types = parse_config(config_path, enable_type_reduction) + required_ops, op_type_impl_filter = parse_config(config_path, enable_type_reduction) - if globally_allowed_types is not None: - globally_allowed_types = _validated_globally_allowed_types(globally_allowed_types) + _process_provider_registrations(ort_root, use_cuda, required_ops, op_type_impl_filter) - _process_provider_registrations(ort_root, use_cuda, required_ops, op_type_usage_manager, globally_allowed_types) - - if op_type_usage_manager is not None: - type_control_cpp_code = op_type_usage_manager.get_cpp_entries() - elif globally_allowed_types is not None: - type_control_cpp_code = ["ORT_SPECIFY_OP_KERNEL_GLOBAL_ALLOWED_TYPES({});".format( - ", ".join(sorted(globally_allowed_types)))] - else: - type_control_cpp_code = [] + type_control_cpp_code = op_type_impl_filter.get_cpp_entries() if op_type_impl_filter is not None else [] _insert_type_control_cpp_code(ort_root, type_control_cpp_code) diff --git a/tools/python/util/ort_format_model/__init__.py b/tools/python/util/ort_format_model/__init__.py index e5ec840d42..5ee341c5f0 100644 --- a/tools/python/util/ort_format_model/__init__.py +++ b/tools/python/util/ort_format_model/__init__.py @@ -12,4 +12,7 @@ sys.path.append(ort_fbs_py_path) from .utils import create_config_from_models # noqa from .ort_model_processor import OrtFormatModelProcessor # noqa -from .operator_type_usage_processors import OperatorTypeUsageManager # noqa +from .operator_type_usage_processors import ( # noqa + GloballyAllowedTypesOpTypeImplFilter, + OpTypeImplFilterInterface, + OperatorTypeUsageManager) diff --git a/tools/python/util/ort_format_model/operator_type_usage_processors.py b/tools/python/util/ort_format_model/operator_type_usage_processors.py index ff9b6d7c3c..6d2b5c84a6 100644 --- a/tools/python/util/ort_format_model/operator_type_usage_processors.py +++ b/tools/python/util/ort_format_model/operator_type_usage_processors.py @@ -2,10 +2,11 @@ # Licensed under the MIT License. import json +import typing import ort_flatbuffers_py.experimental.fbs as fbs from abc import ABC, abstractmethod -from .types import value_name_to_typestr +from .types import FbsTypeInfo, value_name_to_typestr def _create_op_key(domain: str, optype: str): @@ -31,6 +32,26 @@ def _ort_constant_for_domain(domain: str): return domain_to_constant_map[domain] +def _reg_type_to_cpp_type(reg_type: str): + if reg_type == "string": + return "std::string" + return reg_type + + +def _split_reg_types(reg_types_str: str): + ''' + Split on underscores but append "_t" to the previous element. + ''' + tokens = reg_types_str.split("_") + reg_types = [] + for token in tokens: + if token == "t" and len(reg_types) > 0: + reg_types[-1] += "_t" + else: + reg_types += [token] + return reg_types + + class TypeUsageProcessor(ABC): ''' Abstract base class for processors which implement operator specific logic to determine the type or types required. @@ -44,22 +65,25 @@ class TypeUsageProcessor(ABC): def process_node(self, node: fbs.Node, value_name_to_typeinfo: dict): pass - def is_typed_registration_needed(self, type_in_registration): + def is_typed_registration_needed(self, type_in_registration: str, + globally_allowed_types: typing.Optional[typing.Set[str]]): ''' Given the string from a kernel registration, determine if the registration is required or not. :param type_in_registration: Type string from kernel registration + :param globally_allowed_types: Optional set of globally allowed types. If provided, these types take precedence + in determining the required types. :return: True is required. False if not. ''' # Not all operators have typed registrations, so this is optionally implemented by derived classes raise RuntimeError('Did not expect processor for {} to have typed registrations.'.format(self.name)) - @abstractmethod def get_cpp_entry(self): ''' Get the C++ code that specifies this operator's required types. :return: List with any applicable C++ code for this operator's required types. One line per entry. ''' - pass + # Not applicable for some ops, so return no lines by default. + return [] @abstractmethod def to_config_entry(self): @@ -84,14 +108,23 @@ class DefaultTypeUsageProcessor(TypeUsageProcessor): Operator processor which tracks the types used for selected input/s and/or output/s. ''' - def __init__(self, domain: str, optype: str, inputs: [int] = [0], outputs: [int] = []): + def __init__(self, domain: str, optype: str, inputs: [int] = [0], outputs: [int] = [], + required_input_types: typing.Dict[int, typing.Set[str]] = {}, + required_output_types: typing.Dict[int, typing.Set[str]] = {}): ''' Create DefaultTypeUsageProcessor. Types for one or more inputs and/or outputs can be tracked by the processor. The default is to track the types required for input 0, as this is the most common use case in ONNX. + + Required input and output types may be specified. These are only applicable to is_typed_registration_needed(). + If a registration type matches a required type, the typed registration is needed. + There is a separate mechanism for specifying required types from C++ for kernels with untyped registration. + :param domain: Operator domain. :param optype: Operator name. :param inputs: Inputs to track. Zero based index. May be empty. :param outputs: Outputs to track. Zero based index. May be empty. + :param required_input_types: Required input types. May be empty. + :param required_output_types: Required output types. May be empty. ''' super().__init__(domain, optype) self._input_types = {} @@ -106,6 +139,25 @@ class DefaultTypeUsageProcessor(TypeUsageProcessor): if not inputs and not outputs: raise ValueError('At least one input or output must be tracked') + self._required_input_types = required_input_types + self._required_output_types = required_output_types + + def _is_type_enabled(self, reg_type, index, required_types, allowed_type_set): + cpp_type = _reg_type_to_cpp_type(reg_type) + return cpp_type in required_types.get(index, set()) or cpp_type in allowed_type_set + + def is_input_type_enabled(self, reg_type, index, allowed_type_set=None): + '''Whether input type is enabled based on required and allowed types.''' + if allowed_type_set is None: + allowed_type_set = self._input_types[index] + return self._is_type_enabled(reg_type, index, self._required_input_types, allowed_type_set) + + def is_output_type_enabled(self, reg_type, index, allowed_type_set=None): + '''Whether output type is enabled based on required and allowed types.''' + if allowed_type_set is None: + allowed_type_set = self._output_types[index] + return self._is_type_enabled(reg_type, index, self._required_output_types, allowed_type_set) + def process_node(self, node: fbs.Node, value_name_to_typeinfo: dict): for i in self._input_types.keys(): if i >= node.InputsLength(): @@ -128,13 +180,14 @@ class DefaultTypeUsageProcessor(TypeUsageProcessor): type_str = value_name_to_typestr(node.Outputs(o), value_name_to_typeinfo) self._output_types[o].add(type_str) - def is_typed_registration_needed(self, type_in_registration: str): + def is_typed_registration_needed(self, type_in_registration: str, + globally_allowed_types: typing.Optional[typing.Set[str]]): if 0 not in self._input_types.keys(): # currently all standard typed registrations are for input 0. # custom registrations can be handled by operator specific processors (e.g. OneHotProcessor below). raise RuntimeError('Expected typed registration to use type from input 0. Node:{}'.format(self.name)) - return type_in_registration in self._input_types[0] + return self.is_input_type_enabled(type_in_registration, 0, globally_allowed_types) def get_cpp_entry(self): entries = [] @@ -195,8 +248,9 @@ class Output0TypedRegistrationProcessor(DefaultTypeUsageProcessor): # init with tracking of output 0 only. super().__init__(domain, optype, inputs=[], outputs=[0]) - def is_typed_registration_needed(self, type_in_registration: str): - return type_in_registration in self._output_types[0] + def is_typed_registration_needed(self, type_in_registration: str, + globally_allowed_types: typing.Optional[typing.Set[str]]): + return self.is_output_type_enabled(type_in_registration, 0, globally_allowed_types) class OneHotProcessor(TypeUsageProcessor): @@ -212,18 +266,18 @@ class OneHotProcessor(TypeUsageProcessor): type0 = value_name_to_typestr(node.Inputs(0), value_name_to_typeinfo) type1 = value_name_to_typestr(node.Inputs(1), value_name_to_typeinfo) type2 = value_name_to_typestr(node.Inputs(2), value_name_to_typeinfo) - key = '{}_{}_{}'.format(type0, type1, type2) + # types in kernel registration are ordered this way: input (T1), output (T3), depth (T2) + key = (type0, type2, type1) self._triples.add(key) - def is_typed_registration_needed(self, type_in_registration): - # the OneHot registration involves a concatenation of the 3 types involved, in the format we match - # when adding values in process_node - return type_in_registration in self._triples - - def get_cpp_entry(self): - # exclusion is via commenting out the registration entry, so don't need to write any #defines - # to disable type support for the OneHot operator - return None + def is_typed_registration_needed(self, type_in_registration: str, + globally_allowed_types: typing.Optional[typing.Set[str]]): + # the OneHot registration involves a concatenation of the 3 types involved + reg_types = tuple([_reg_type_to_cpp_type(reg_type) for reg_type in _split_reg_types(type_in_registration)]) + if globally_allowed_types is not None: + return all(reg_type in globally_allowed_types for reg_type in reg_types) + else: + return reg_types in self._triples def to_config_entry(self): if not self._triples: @@ -237,7 +291,7 @@ class OneHotProcessor(TypeUsageProcessor): self._triples.clear() aggregate_info = json.loads(entry) if 'custom' in aggregate_info: - self._triples = set(aggregate_info['custom']) + self._triples = set([tuple(triple) for triple in aggregate_info['custom']]) def _create_operator_type_usage_processors(): @@ -267,25 +321,33 @@ def _create_operator_type_usage_processors(): # - Implementation does not have any significant type specific code: # ai.onnx: Concat, Flatten, Not, QLinearConv, Reshape, Shape, Squeeze, Unsqueeze # - default_processor_onnx_ops = ['Abs', 'Add', 'ArgMax', 'ArgMin', 'AveragePool', + default_processor_onnx_ops = ['Abs', 'ArgMax', 'ArgMin', 'AveragePool', 'BatchNormalization', 'BitShift', 'Ceil', 'Clip', 'Conv', 'CumSum', - 'DequantizeLinear', 'Div', - 'Equal', 'Exp', 'Expand', + 'DequantizeLinear', + 'Exp', 'Expand', 'Floor', - 'Gemm', 'Greater', - 'IsNaN' - 'Less', 'Log', 'LogSoftmax', 'LpNormalization', - 'MatMul', 'Max', 'Min', 'Mul', + 'Gemm', + 'IsNaN', + 'Log', 'LogSoftmax', 'LpNormalization', + 'MatMul', 'Max', 'Min', 'Neg', 'NonMaxSuppression', 'NonZero', 'Pad', 'Range', 'Reciprocal', 'ReduceL1', 'ReduceL2', 'ReduceLogSum', 'ReduceLogSumExp', 'ReduceMax', 'ReduceMean', 'ReduceMin', 'ReduceProd', 'ReduceSum', 'ReduceSumSquare', 'Relu', 'Resize', 'RoiAlign', 'Round', - 'Sigmoid', 'Sin', 'Softmax', 'Split', 'Sqrt', 'Sub', + 'Sigmoid', 'Sin', 'Softmax', 'Split', 'Sqrt', 'Tanh', 'Tile', 'TopK', 'Transpose', 'Where'] + default_processor_onnx_ops_requiring_int64_for_input_0 = ['Add', + 'Div', + 'Equal', + 'Greater', + 'Less', + 'Mul', + 'Sub'] + internal_ops = ['QLinearAdd', 'QLinearMul'] # TODO - review and add ML ops as needed @@ -300,6 +362,8 @@ def _create_operator_type_usage_processors(): default_processor_onnxml_ops = [] [add(DefaultTypeUsageProcessor('ai.onnx', op)) for op in default_processor_onnx_ops] + [add(DefaultTypeUsageProcessor('ai.onnx', op, required_input_types={0: {"int64_t"}})) + for op in default_processor_onnx_ops_requiring_int64_for_input_0] [add(DefaultTypeUsageProcessor('ai.onnx.ml', op)) for op in default_processor_onnxml_ops] [add(DefaultTypeUsageProcessor('com.microsoft', op)) for op in internal_ops] @@ -335,6 +399,30 @@ def _create_operator_type_usage_processors(): return operator_processors +class OpTypeImplFilterInterface(ABC): + ''' + Class that filters operator implementations based on type. + ''' + @abstractmethod + def is_typed_registration_needed(self, domain: str, optype: str, type_registration_str: str): + ''' + Given the string from a kernel registration, determine if the registration is required or not. + :param domain: Operator domain. + :param optype: Operator type. + :param type_registration_str: Type string from kernel registration + :return: True is required. False if not. + ''' + pass + + @abstractmethod + def get_cpp_entries(self): + ''' + Get the C++ code that specifies the operator types to enable. + :return: List of strings. One line of C++ code per entry. + ''' + pass + + class OperatorTypeUsageManager: ''' Class to manage the operator type usage processors. @@ -371,32 +459,6 @@ class OperatorTypeUsageManager: if op_processor: op_processor.process_node(node, value_name_to_typeinfo) - def is_typed_registration_needed(self, domain: str, optype: str, type_registration_str: str): - ''' - Given the string from a kernel registration, determine if the registration is required or not. - :param domain: Operator domain. - :param optype: Operator type. - :param type_registration_str: Type string from kernel registration - :return: True is required. False if not. - ''' - needed = True # we keep the registration unless the per-operator processor says not to - key = _create_op_key(domain, optype) - if key in self._operator_processors: - needed = self._operator_processors[key].is_typed_registration_needed(type_registration_str) - - return needed - - def get_cpp_entries(self): - ''' - Get the C++ code that define the lists of types to enable for the operators we have type info for. - :return: List of strings. One line of C++ code per entry. - ''' - entries = [] - for key in sorted(self._operator_processors.keys()): - entries.extend(self._operator_processors[key].get_cpp_entry()) - - return entries - def get_config_entry(self, domain: str, optype: str): ''' Get the config entry specifying the types for this operator. @@ -438,3 +500,61 @@ class OperatorTypeUsageManager: # same values back self._operator_processors[key].from_config_entry(entry) assert(entry == self._operator_processors[key].to_config_entry()) + + class _OpTypeImplFilter(OpTypeImplFilterInterface): + def __init__(self, manager): + self._manager = manager + + def is_typed_registration_needed(self, domain: str, optype: str, type_registration_str: str): + needed = True # we keep the registration unless the per-operator processor says not to + key = _create_op_key(domain, optype) + if key in self._manager._operator_processors: + needed = self._manager._operator_processors[key].is_typed_registration_needed( + type_in_registration=type_registration_str, globally_allowed_types=None) + + return needed + + def get_cpp_entries(self): + entries = [] + for key in sorted(self._manager._operator_processors.keys()): + entries.extend(self._manager._operator_processors[key].get_cpp_entry()) + + return entries + + def make_op_type_impl_filter(self): + ''' + Creates an OpTypeImplFilterInterface instance from this manager. + Filtering uses the manager's operator type usage processor state. + ''' + return OperatorTypeUsageManager._OpTypeImplFilter(self) + + +class GloballyAllowedTypesOpTypeImplFilter(OpTypeImplFilterInterface): + ''' + Operator implementation filter which uses globally allowed types. + ''' + _valid_allowed_types = set(FbsTypeInfo.tensordatatype_to_string.values()) + + def __init__(self, globally_allowed_types: typing.Set[str]): + self._operator_processors = _create_operator_type_usage_processors() + + if not globally_allowed_types.issubset(self._valid_allowed_types): + raise ValueError("Globally allowed types must all be valid. Invalid types: {}" + .format(sorted(globally_allowed_types - self._valid_allowed_types))) + + self._globally_allowed_types = globally_allowed_types + + def is_typed_registration_needed(self, domain: str, optype: str, type_registration_str: str): + key = _create_op_key(domain, optype) + if key in self._operator_processors: + needed = self._operator_processors[key].is_typed_registration_needed( + type_in_registration=type_registration_str, + globally_allowed_types=self._globally_allowed_types) + else: + needed = _reg_type_to_cpp_type(type_registration_str) in self._globally_allowed_types + + return needed + + def get_cpp_entries(self): + return ["ORT_SPECIFY_OP_KERNEL_GLOBAL_ALLOWED_TYPES({});".format( + ", ".join(sorted(self._globally_allowed_types)))] diff --git a/tools/python/util/ort_format_model/types.py b/tools/python/util/ort_format_model/types.py index a2f69b3735..76e69615b4 100644 --- a/tools/python/util/ort_format_model/types.py +++ b/tools/python/util/ort_format_model/types.py @@ -20,8 +20,8 @@ class FbsTypeInfo: fbs.TensorDataType.TensorDataType.DOUBLE: 'double', fbs.TensorDataType.TensorDataType.UINT32: 'uint32_t', fbs.TensorDataType.TensorDataType.UINT64: 'uint64_t', - fbs.TensorDataType.TensorDataType.COMPLEX64: 'complex64 is not supported', - fbs.TensorDataType.TensorDataType.COMPLEX128: 'complex128 is not supported', + # fbs.TensorDataType.TensorDataType.COMPLEX64: 'complex64 is not supported', + # fbs.TensorDataType.TensorDataType.COMPLEX128: 'complex128 is not supported', fbs.TensorDataType.TensorDataType.BFLOAT16: 'BFloat16' } diff --git a/tools/python/util/reduced_build_config_parser.py b/tools/python/util/reduced_build_config_parser.py index f49ae637b2..420111ee01 100644 --- a/tools/python/util/reduced_build_config_parser.py +++ b/tools/python/util/reduced_build_config_parser.py @@ -7,21 +7,24 @@ import os try: import flatbuffers # noqa have_flatbuffers = True - from .ort_format_model import OperatorTypeUsageManager # noqa + from .ort_format_model import GloballyAllowedTypesOpTypeImplFilter, OperatorTypeUsageManager # noqa except ImportError: have_flatbuffers = False def parse_config(config_file: str, enable_type_reduction: bool = False): ''' - Parse the configuration file and return the required operators dictionary, and possibly either an - OperatorTypeUsageManager or a globally allowed types list. + Parse the configuration file and return the required operators dictionary and an + OpTypeImplFilterInterface instance. Configuration file lines can do the following: 1. specify required operators 2. specify globally allowed types for all operators + 3. specify what it means for no required operators to be specified - The basic format for specifying required operators (1) is `domain;opset;op1,op2...` + 1. Specifying required operators + + The basic format for specifying required operators is `domain;opset;op1,op2...` e.g. `ai.onnx;11;Add,Cast,Clip,... If the configuration file is generated from ORT format models it may optionally contain JSON for per-operator @@ -41,50 +44,72 @@ def parse_config(config_file: str, enable_type_reduction: bool = False): `{"inputs": {"0": ["float", "int32_t"], "1": ["int32_t"]}}` Finally some operators do non-standard things and store their type information under a 'custom' key. - ai.onnx.OneHot is an example of this, where 3 type names from the inputs are combined into a string. - `{"custom": ["float_int64_t_int64_t", "int64_t_string_int64_t"]}` + ai.onnx.OneHot is an example of this, where the three input types are combined into a triple. + `{"custom": [["float", "int64_t", "int64_t"], ["int64_t", "std::string", "int64_t"]]}` + + 2. Specifying globally allowed types for all operators + + The format for specifying globally allowed types for all operators is: + `!globally_allowed_types;T0,T1,...` - The format for specifying globally allowed types for all operators (2) is `!globally_allowed_types;T0,T1,...` Ti should be a C++ scalar type supported by ONNX and ORT. At most one globally allowed types specification is allowed. Specifying per-operator type information and specifying globally allowed types are mutually exclusive - it is an error to specify both. + 3. Specify what it means for no required operators to be specified + + By default, if no required operators are specified, NO operators are required. + + With the following line, if no required operators are specified, ALL operators are required: + `!no_ops_specified_means_all_ops_are_required` + :param config_file: Configuration file to parse :param enable_type_reduction: Set to True to use the type information in the config. If False the type information will be ignored. - If the flatbuffers module is unavailable op-specific type information will be ignored - as the per-op usage of the type information is via OperatorTypeUsageManager which has - a dependency on the ORT flatbuffers python schema. - :return: required_ops, op_type_usage_manager, globally_allowed_types: - Dictionary of domain:opset:[ops] for required operators - OperatorTypeUsageManager manager with operator specific type usage information if available. None if - type reduction was disabled, per-op type reduction was not specified, or the flatbuffers module is not - available. - List of globally allowed types. None if type reduction was disabled, or globally allowed types were not - specified. - At most one of the op_type_usage_manager and globally_allowed_types tuple elements will not be None. + If the flatbuffers module is unavailable type information will be ignored as the + type-based filtering has a dependency on the ORT flatbuffers schema. + :return: required_ops: Dictionary of domain:opset:[ops] for required operators. If None, all operators are + required. + op_type_impl_filter: OpTypeImplFilterInterface instance if type reduction is enabled, the flatbuffers + module is available, and type reduction information is present. None otherwise. ''' if not os.path.isfile(config_file): raise ValueError('Configuration file {} does not exist'.format(config_file)) + # only enable type reduction when flatbuffers is available + enable_type_reduction = enable_type_reduction and have_flatbuffers + required_ops = {} - op_type_usage_manager = OperatorTypeUsageManager() if enable_type_reduction and have_flatbuffers else None + no_ops_specified_means_all_ops_are_required = False + op_type_usage_manager = OperatorTypeUsageManager() if enable_type_reduction else None has_op_type_reduction_info = False globally_allowed_types = None + def process_non_op_line(line): + if not line or line.startswith("#"): # skip empty lines and comments + return True + + if line.startswith("!globally_allowed_types;"): # handle globally allowed types + if enable_type_reduction: + nonlocal globally_allowed_types + if globally_allowed_types is not None: + raise RuntimeError("Globally allowed types were already specified.") + globally_allowed_types = set(segment.strip() for segment in line.split(';')[1].split(',')) + return True + + if line == "!no_ops_specified_means_all_ops_are_required": # handle all ops required line + nonlocal no_ops_specified_means_all_ops_are_required + no_ops_specified_means_all_ops_are_required = True + return True + + return False + with open(config_file, 'r') as config: for line in [orig_line.strip() for orig_line in config.readlines()]: - if not line or line.startswith("#"): # skip empty lines and comments - continue - - if line.startswith("!globally_allowed_types;"): # handle globally allowed types - if enable_type_reduction: - if globally_allowed_types is not None: - raise RuntimeError("Globally allowed types were already specified.") - globally_allowed_types = [segment.strip() for segment in line.split(';')[1].split(',')] + if process_non_op_line(line): continue domain, opset_str, operators_str = [segment.strip() for segment in line.split(';')] @@ -151,6 +176,10 @@ def parse_config(config_file: str, enable_type_reduction: bool = False): else: required_ops[domain][opset].update(operators) + if len(required_ops) == 0 and no_ops_specified_means_all_ops_are_required: + required_ops = None + + op_type_impl_filter = None if enable_type_reduction: if not has_op_type_reduction_info: op_type_usage_manager = None @@ -158,4 +187,9 @@ def parse_config(config_file: str, enable_type_reduction: bool = False): raise RuntimeError( "Specifying globally allowed types and per-op type reduction info together is unsupported.") - return required_ops, op_type_usage_manager, globally_allowed_types + if globally_allowed_types is not None: + op_type_impl_filter = GloballyAllowedTypesOpTypeImplFilter(globally_allowed_types) + elif op_type_usage_manager is not None: + op_type_impl_filter = op_type_usage_manager.make_op_type_impl_filter() + + return required_ops, op_type_impl_filter