From 8c425a106e71bfcc3ac5cedf6d594c35e40f756a Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Wed, 23 Mar 2022 06:47:37 -0700 Subject: [PATCH 1/8] [TRT] Move ops from predicates to DFPattern-based to cater transition to Collage --- python/tvm/relay/op/contrib/tensorrt.py | 213 +++++++++++++---- python/tvm/relay/transform/transform.py | 18 +- src/relay/transforms/unmerge_composites.cc | 254 +++++++++++++++++++++ 3 files changed, 435 insertions(+), 50 deletions(-) create mode 100644 src/relay/transforms/unmerge_composites.cc diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index 3bd737e6e0fd..986489bae29d 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -24,8 +24,10 @@ from tvm.ir import Op from tvm.relay import transform from tvm.relay.build_module import bind_params_by_name +from tvm.relay.dataflow_pattern import is_op, wildcard from tvm.relay.expr import Call, Constant, GlobalVar, Tuple, TupleGetItem, Var from tvm.relay.expr_functor import ExprMutator, ExprVisitor +from tvm.relay.op.contrib.register import register_pattern_table logger = logging.getLogger("TensorRT") supported_types = ["float32", "float16"] @@ -39,7 +41,8 @@ def is_supported_trt_dtype(args): True if supported, False if not. """ if not all([x.checked_type.dtype in supported_types for x in args]): - logger.info("Only float32 and float16 inputs are supported for TensorRT BYOC.") + logger.info( + "Only float32 and float16 inputs are supported for TensorRT BYOC.") return False return True @@ -51,7 +54,8 @@ def is_tensorrt_runtime_enabled(): ret: bool True if present, False if not. """ - check_enabled = tvm.get_global_func("relay.op.is_tensorrt_runtime_enabled", True) + check_enabled = tvm.get_global_func( + "relay.op.is_tensorrt_runtime_enabled", True) if check_enabled: return check_enabled() return False @@ -150,7 +154,8 @@ def partition_for_tensorrt( assert isinstance(version, tuple) and len(version) == 3 config["tensorrt_version"] = version else: - linked_version = tuple(tvm.get_global_func("relay.op.get_tensorrt_version")()) + linked_version = tuple(tvm.get_global_func( + "relay.op.get_tensorrt_version")()) if not linked_version: logger.warning( "TVM was not built against TensorRT and no version was provided to " @@ -175,14 +180,18 @@ def partition_for_tensorrt( } ), transform.FoldConstant(), + transform.MergeComposite(pattern_table()), + # transform.AnnotateTarget(["tensorrt"], include_non_call_ops=True), transform.AnnotateTarget("tensorrt"), transform.MergeCompilerRegions(), transform.PartitionGraph(), + transform.UnmergeComposites(), transform.InferType(), ] ) with tvm.transform.PassContext(opt_level=3, config={"relay.ext.tensorrt.options": config}): mod = seq(mod) + print(mod) mod = prune_tensorrt_subgraphs(mod) return mod, config @@ -211,13 +220,14 @@ def check_dynamism(args, op_name): elif isinstance(arg, Tuple): return check_dynamism(arg.fields, op_name) else: - logger.info("Arg not supported in TensorRT for %s with type %s", op_name, type(arg)) + logger.info( + "Arg not supported in TensorRT for %s with type %s", op_name, type(arg)) return True return False def _register_external_op_helper_with_checker(op_name, checker): - @tvm.ir.register_op_attr(op_name, "target.tensorrt") + @ tvm.ir.register_op_attr(op_name, "target.tensorrt") def _func_wrapper(expr): attrs, args = expr.attrs, expr.args # ops with dynamic shapes are offloaded to VM @@ -237,7 +247,8 @@ def _func_wrapper(expr): # have been excluded because they occur in PT MaskRCNN model. The long term solution is # to switch to explicit batch mode after performance regressions are solved. if all( - [list(map(int, shape)) in [[300, 64, 7, 7], [300, 1, 1, 1]] for shape in shapes] + [list(map(int, shape)) in [[300, 64, 7, 7], [300, 1, 1, 1]] + for shape in shapes] ): return False return checker(attrs, args, op_name) @@ -314,7 +325,8 @@ def trt_version_annotate_fn(version): def _func_wrapper(attrs, args, op_name): if get_tensorrt_version() < version: logger.info( - "%s: requires TensorRT version %s or higher.", op_name, ".".join(map(str, version)) + "%s: requires TensorRT version %s or higher.", op_name, ".".join( + map(str, version)) ) return False return True @@ -322,12 +334,18 @@ def _func_wrapper(attrs, args, op_name): return _func_wrapper -_register_external_op_helper_with_checker("nn.leaky_relu", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker("sin", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker("cos", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker("atan", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker("ceil", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker("erf", trt_version_annotate_fn((7, 0, 0))) +_register_external_op_helper_with_checker( + "nn.leaky_relu", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker( + "sin", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker( + "cos", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker( + "atan", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker( + "ceil", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker( + "erf", trt_version_annotate_fn((7, 0, 0))) @_register_external_dynamic_check_func("add") @@ -338,7 +356,8 @@ def add_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False shapes = [ - [int(x) if not isinstance(x, tvm.tir.expr.Any) else -1 for x in arg.checked_type.shape] + [int(x) if not isinstance(x, tvm.tir.expr.Any) + else -1 for x in arg.checked_type.shape] for arg in args ] @@ -368,13 +387,15 @@ def batch_norm_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if len(args[0].checked_type.shape) == 5 and get_tensorrt_version() < (6, 0, 1): - logger.info("nn.batch_norm: TensorRT 6.0.1 or higher is required for rank 5 inputs.") + logger.info( + "nn.batch_norm: TensorRT 6.0.1 or higher is required for rank 5 inputs.") return False if len(args[0].checked_type.shape) > 5: logger.info("nn.batch_norm: Input rank must be 5 or less.") return False if int(attrs.axis) not in (1, 3): - logger.info("nn.batch_norm: axis is %d but must be 1 or 3.", int(attrs.axis)) + logger.info("nn.batch_norm: axis is %d but must be 1 or 3.", + int(attrs.axis)) return False return True @@ -400,33 +421,104 @@ def conv1d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCW": - logger.info("nn.conv1d: data_layout is %s but must be NCW.", attrs.data_layout) + logger.info("nn.conv1d: data_layout is %s but must be NCW.", + attrs.data_layout) return False if attrs.kernel_layout != "OIW": - logger.info("nn.conv1d: kernel_layout is %s but must be OIW.", attrs.kernel_layout) + logger.info("nn.conv1d: kernel_layout is %s but must be OIW.", + attrs.kernel_layout) return False return True -@_register_external_dynamic_check_func("nn.conv2d") -def conv2d_annotate_fn(expr): # pylint: disable=unused-variable - """Check if nn.conv2d is supported by TensorRT.""" +def conv2d_pattern(): + conv2d = is_op("nn.conv2d")(wildcard(), wildcard()) + return conv2d - attrs, args = expr.attrs, expr.args + +def check_conv2d(pattern): + attrs, args = pattern.attrs, pattern.args if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCHW": - logger.info("nn.conv2d: data_layout is %s but must be NCHW.", attrs.data_layout) + logger.info( + "nn.conv2d: data_layout is %s but must be NCHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIHW": - logger.info("nn.conv2d: kernel_layout is %s but must be OIHW.", attrs.kernel_layout) + logger.info( + "nn.conv2d: kernel_layout is %s but must be OIHW.", attrs.kernel_layout) return False if attrs.out_layout and attrs.out_layout != "NCHW": - logger.info("nn.conv2d: out_layout is %s but must be NCHW.", attrs.out_layout) + logger.info( + "nn.conv2d: out_layout is %s but must be NCHW.", attrs.out_layout) + return False + return True + + +def add_pattern(): + pat = is_op("add")(wildcard(), wildcard()) + return pat + + +def check_add(pattern): # pylint: disable=unused-variable + """Check if add is supported by TensorRT.""" + + args = pattern.args + if not is_supported_trt_dtype(args): + return False + shapes = [ + [int(x) if not isinstance(x, tvm.tir.expr.Any) + else -1 for x in arg.checked_type.shape] + for arg in args + ] + + # Scalars require explicit batch mode. + if get_tensorrt_use_implicit_batch_mode() and any([len(shape) < 1 for shape in shapes]): + return False + + if ( + not get_tensorrt_use_implicit_batch_mode() + and (isinstance(args[0], Constant) or isinstance(args[1], Constant)) + and len(shapes[0]) > 0 + and len(shapes[1]) > 0 + and shapes[0][0] == shapes[1][0] + and shapes[0][0] != 1 + and (len(shapes[0]) > 3 or len(shapes[1]) > 3) + ): + logger.info("add: bug in TRT with adding batched constants.") return False return True +def squeeze_pattern(): + """Create the pattern for squeeze.""" + return is_op("squeeze")(wildcard()) + + +def check_squeeze(pattern): + attrs, args = pattern.attrs, pattern.args + if not is_supported_trt_dtype(args): + return False + if not attrs.axis: + logger.info("squeeze: must explicitly set axis.") + return False + if get_tensorrt_use_implicit_batch_mode() and any([axis == 0 for axis in map(int, attrs.axis)]): + logger.info("squeeze: can't modify batch dimension.") + return False + return True + + +def batch_mul_pattern(): + """Create the pattern for squeeze.""" + return is_op("nn.batch_matmul")(wildcard(), wildcard()) + + +@register_pattern_table("tensorrt") +def pattern_table(): + return [("tensorrt.nn.conv2d", conv2d_pattern(), check_conv2d), ("tensorrt.squeeze", squeeze_pattern(), check_squeeze), + ("tensorrt.add", add_pattern(), check_add), ("tensorrt.nn.batch_matmul", batch_mul_pattern())] + + @_register_external_dynamic_check_func("nn.dense") def dense_annotate_fn(expr): # pylint: disable=unused-variable """Check if dense is supported by TensorRT.""" @@ -437,7 +529,8 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable input_rank = len(args[0].checked_type.shape) weight_rank = len(args[1].checked_type.shape) if input_rank not in (2, 3, 4): - logger.info("nn.dense: input has rank %d but must be 2, 3 or 4.", input_rank) + logger.info( + "nn.dense: input has rank %d but must be 2, 3 or 4.", input_rank) return False if weight_rank != 2: logger.info("nn.dense: weight has rank %d but must be 2.", weight_rank) @@ -445,7 +538,7 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable return True -@_register_external_dynamic_check_func("nn.batch_matmul") +@_register_external_dynamic_check_func("tensorrt.nn.batch_matmul") def batch_matmul_annotate_fn(expr): """Check if dense is supported by TensorRT.""" @@ -482,7 +575,8 @@ def bias_add_annotate_fn(expr): # pylint: disable=unused-variable return False input_rank = len(args[0].checked_type.shape) if input_rank not in (2, 3, 4): - logger.info("nn.bias_add: input rank is %d but must be 2, 3 or 4.", input_rank) + logger.info( + "nn.bias_add: input rank is %d but must be 2, 3 or 4.", input_rank) return False return True @@ -495,10 +589,12 @@ def max_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info("nn.max_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info( + "nn.max_pool2d: layout is %s but must be NCHW.", attrs.layout) return False if attrs.ceil_mode and get_tensorrt_version() < (5, 1, 5): - logger.info("nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") + logger.info( + "nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") return False return True @@ -511,7 +607,8 @@ def avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info("nn.avg_pool2d: layout is %d but must be NCHW.", attrs.layout) + logger.info( + "nn.avg_pool2d: layout is %d but must be NCHW.", attrs.layout) return False if ( attrs.count_include_pad @@ -527,7 +624,8 @@ def avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.ceil_mode and get_tensorrt_version() < (5, 1, 5): - logger.info("nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") + logger.info( + "nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") return False return True @@ -540,7 +638,8 @@ def global_max_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info("nn.global_max_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info( + "nn.global_max_pool2d: layout is %s but must be NCHW.", attrs.layout) return False return True @@ -553,7 +652,8 @@ def global_avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info("nn.global_avg_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info( + "nn.global_avg_pool2d: layout is %s but must be NCHW.", attrs.layout) return False return True @@ -593,7 +693,8 @@ def concatenate_annotate_fn(expr): # pylint: disable=unused-variable attrs, args = expr.attrs, expr.args if any([x.dtype not in supported_types for x in args[0].checked_type.fields]): - logger.info("Only float16 and float32 inputs are supported for TensorRT.") + logger.info( + "Only float16 and float32 inputs are supported for TensorRT.") if not get_tensorrt_use_implicit_batch_mode(): return True if int(attrs.axis) == 0: @@ -602,7 +703,8 @@ def concatenate_annotate_fn(expr): # pylint: disable=unused-variable if isinstance(args[0], Tuple): for tuple_input in args[0].fields: if isinstance(tuple_input, Constant): - logger.info("concatenate: can't concatenate tensors with constants.") + logger.info( + "concatenate: can't concatenate tensors with constants.") return False return True @@ -628,7 +730,8 @@ def conv2d_transpose_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCHW": - logger.info("nn.conv2d_transpose: data_layout is %s but must be NCHW.", attrs.data_layout) + logger.info( + "nn.conv2d_transpose: data_layout is %s but must be NCHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIHW": logger.info( @@ -636,7 +739,8 @@ def conv2d_transpose_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.out_layout and attrs.out_layout != "NCHW": - logger.info("nn.conv2d_transpose: out_layout is %s but must be NCHW.", attrs.out_layout) + logger.info( + "nn.conv2d_transpose: out_layout is %s but must be NCHW.", attrs.out_layout) return False if attrs.dilation and any([rate != 1 for rate in map(int, attrs.dilation)]): logger.info("nn.conv2d_transpose: dilation rate must be 1.") @@ -725,7 +829,8 @@ def reshape_annotate_fn(expr): # pylint: disable=unused-variable # Resolve -1. for i, value in enumerate(new_shape): if value == -1: - new_shape[i] = original_volume // np.prod([x for x in new_shape if x != -1]) + new_shape[i] = original_volume // np.prod( + [x for x in new_shape if x != -1]) # Remove batch dimension and see if volumes match if shape[0] != new_shape[0]: logger.info("reshape: can't modify batch dimension.") @@ -744,7 +849,8 @@ def pad_annotate_fn(expr): # pylint: disable=unused-variable assert isinstance(pad_value, relay.Constant) pad_value = pad_value.data.numpy().item() if attrs.pad_mode != "constant": - logger.info("nn.pad: pad mode is %s but must be constant.", attrs.pad_mode) + logger.info("nn.pad: pad mode is %s but must be constant.", + attrs.pad_mode) return False if pad_value > 0.0: logger.info("nn.pad: pad value is %f but must be 0.0.", pad_value) @@ -771,7 +877,8 @@ def strided_slice_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((5, 1, 5))(attrs, args, "strided_slice"): return False if get_tensorrt_use_implicit_batch_mode(): - batch_dim_begin_modified = attrs.begin[0] is not None and int(attrs.begin[0]) != 0 + batch_dim_begin_modified = attrs.begin[0] is not None and int( + attrs.begin[0]) != 0 batch_dim_end_modified = ( attrs.end[0] is not None and int(attrs.end[0]) != -1 @@ -844,13 +951,16 @@ def conv3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.conv3d"): return False if attrs.data_layout != "NCDHW": - logger.info("nn.conv3d: data_layout is %s but must be NCDHW.", attrs.data_layout) + logger.info( + "nn.conv3d: data_layout is %s but must be NCDHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIDHW": - logger.info("nn.conv3d: kernel_layout is %s but must be OIDHW.", attrs.kernel_layout) + logger.info( + "nn.conv3d: kernel_layout is %s but must be OIDHW.", attrs.kernel_layout) return False if attrs.out_layout and attrs.out_layout != "NCDHW": - logger.info("nn.conv3d: out_layout is %s but must be NCDHW.", attrs.out_layout) + logger.info( + "nn.conv3d: out_layout is %s but must be NCDHW.", attrs.out_layout) return False return True @@ -865,7 +975,8 @@ def max_pool_3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.max_pool3d"): return False if attrs.layout != "NCDHW": - logger.info("nn.max_pool3d: layout is %s but must be NCDHW.", attrs.layout) + logger.info( + "nn.max_pool3d: layout is %s but must be NCDHW.", attrs.layout) return False return True @@ -880,7 +991,8 @@ def avg_pool_3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.avg_pool3d"): return False if attrs.layout != "NCDHW": - logger.info("nn.avg_pool3d: layout is %s but must be NCDHW.", attrs.layout) + logger.info( + "nn.avg_pool3d: layout is %s but must be NCDHW.", attrs.layout) return False return True @@ -895,7 +1007,8 @@ def conv3d_transpose_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.conv3d_transpose"): return False if attrs.data_layout != "NCDHW": - logger.info("nn.conv3d_transpose: data_layout is %s but must be NCDHW.", attrs.data_layout) + logger.info( + "nn.conv3d_transpose: data_layout is %s but must be NCDHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIDHW": logger.info( @@ -903,7 +1016,8 @@ def conv3d_transpose_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.out_layout and attrs.out_layout != "NCDHW": - logger.info("nn.conv3d_transpose: out_layout is %s but must be NCDHW.", attrs.out_layout) + logger.info( + "nn.conv3d_transpose: out_layout is %s but must be NCDHW.", attrs.out_layout) return False if attrs.dilation and any([rate != 1 for rate in map(int, attrs.dilation)]): logger.info("nn.conv3d_transpose: dilation rate must be 1.") @@ -1033,7 +1147,8 @@ def visit_call(self, call): subgraphs_to_remove.append(name) # Create new pruned module new_mod = tvm.IRModule(mod.functions, mod.type_definitions) - new_mod["main"] = SubgraphRemover(subgraphs_to_remove, mod, new_mod).visit(mod["main"]) + new_mod["main"] = SubgraphRemover( + subgraphs_to_remove, mod, new_mod).visit(mod["main"]) new_mod = transform.RemoveUnusedFunctions()(new_mod) return new_mod diff --git a/python/tvm/relay/transform/transform.py b/python/tvm/relay/transform/transform.py index 99c61c5bd96f..ac0f61a189b5 100644 --- a/python/tvm/relay/transform/transform.py +++ b/python/tvm/relay/transform/transform.py @@ -543,7 +543,10 @@ def MergeComposite(pattern_table): for tup in pattern_table: if len(tup) == 2: pattern_name, pattern = tup - check = lambda extract: True + + def check(extract): + return True + elif len(tup) == 3: pattern_name, pattern, check = tup @@ -789,6 +792,19 @@ def Inline(): return _ffi_api.Inline() +def UnmergeComposites(): + """Perform inlining on the given Relay IR module. The global functions that + are marked as `inline` should be always inlined. A cost model will be + needed in the future to decide if it is profitable to inline the function. + + Returns + ------- + ret: tvm.transform.Pass + The registered pass that performs inlining for a Relay IR module. + """ + return _ffi_api.UnmergeComposites() + + def gradient(expr, mod=None, mode="higher_order"): """ Transform the input function, diff --git a/src/relay/transforms/unmerge_composites.cc b/src/relay/transforms/unmerge_composites.cc new file mode 100644 index 000000000000..72c3ec6978ab --- /dev/null +++ b/src/relay/transforms/unmerge_composites.cc @@ -0,0 +1,254 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +/*! + * \file src/relay/transforms/unmerge_composites.cc + * \brief Unmerges composite functions + */ +#include +#include +#include + +#include "../analysis/call_graph.h" +#include "../op/call/call.h" + +using namespace tvm::runtime; + +namespace tvm { + +namespace relay { + +class Unmerger : ExprMutator { + public: + explicit Unmerger(CallGraphEntry* cur_node, CallGraphNode* call_graph) + : cur_node_(cur_node), call_graph_(call_graph) {} + + Expr VisitExpr_(const CallNode* call_node) final { + // We can work with calls in both pre- and post-lowered form. + Call vanilla_call = GetAnyCall(call_node); + VLOG(1) << "Vanilla call " << vanilla_call->op << std::endl; + VLOG(1) << "Vanilla call " << vanilla_call->op->checked_type_ << std::endl; + // VLOG(1) << "Vanilla attrs " << vanilla_call->attrs << std::endl; + const auto* global_var_node = vanilla_call->op.as(); + const auto* function_var_node = vanilla_call->op.as(); + // const auto* function__node = vanilla_call->op.as(); + + if (global_var_node) { + VLOG(1) << "Global Var node"; + } + + if (function_var_node) { + VLOG(1) << "For existing function "; + Function gv = GetRef(function_var_node); + const auto* fn = gv.as(); + ICHECK(fn) << "Expected to work on a Relay function."; + // auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, fn->attrs); + + Array new_args; + new_args.reserve(vanilla_call->args.size()); + for (auto arg : vanilla_call->args) { + new_args.push_back(VisitExpr(arg)); + } + + Map bind_map; + for (size_t i = 0; i < new_args.size(); i++) { + bind_map.Set(fn->params[i], new_args[i]); + } + + auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, {}); + VLOG(1) << "Params :" << func->params; + VLOG(1) << "Ret type :" << func->ret_type; + VLOG(1) << "Type Params :" << func->type_params; + VLOG(1) << "Attrs :" << func->attrs; + VLOG(1) << "Func body:" << func->body; + + return Bind(func->body, bind_map); + + // return func; + // return func->body; + // VLOG(1) << "gv " << func; + // return func->body; + // auto base_func = call_graph_->GetGlobalFunction(global); + } + // ICHECK(function_var_node); + // VLOG(1) << std::endl << "function : " << function__node << std::endl; + // VLOG(1) << "Function var node : " << function_var_node->attrs; + // if(function_var_node) { + // VLOG(1) << "As function var node"; + // } + // if (global_var_node) { + // ICHECK(function_var_node); + // ICHECK(global_var_node); + // GlobalVar gv = GetRef(global_var_node); + // Function gv = GetRef(function_var_node); + + // auto* cg_node = (*call_graph_)[gv->name_hint]; + // // if (CanInline(cg_node)) { + // Array new_args; + // new_args.reserve(vanilla_call->args.size()); + // for (auto arg : vanilla_call->args) { + // new_args.push_back(VisitExpr(arg)); + // } + // cur_node_->RemoveCallTo(gv); + // return MakeNewExpr(gv, new_args, GetRef(call_node)); + // } + // else: fallthrough + // } + // else: fallthrough + + // If not calling a global function then nothing to inline. + return ExprMutator::VisitExpr_(call_node); + } + + Expr VisitExpr_(const GlobalVarNode* gvn) final { + GlobalVar gv = GetRef(gvn); + auto* cg_node = (*call_graph_)[gv->name_hint]; + if (CanInline(cg_node)) { + cur_node_->RemoveCallTo(gv); + return MakeNewExpr(gv, {}, GetRef(gvn)); + } + return ExprMutator::VisitExpr_(gvn); + } + + Function Unmerge(const Function& func) { + return WithFields(func, func->params, VisitExpr(func->body)); + } + + private: + bool CanInline(const CallGraphEntry* cg_node) { + // The node must be a leaf node and it cannot be recursive. + return true; + if (!cg_node->empty() || cg_node->IsRecursive()) return false; + + auto base_func = call_graph_->GetGlobalFunction(cg_node->GetGlobalVar()); + const auto* function_node = base_func.as(); + if (!function_node) { + // Can't inline PrimFuncs! + return false; + } + // The body of a global functions must be defined. + if (!function_node->body.defined()) return false; + + // The function must be annotated with the inline attribute. + // (Note that external functions do not have this attribute!) + // if (!function_node->HasNonzeroAttr(attr::kInline)) return false; + if (!function_node->HasNonzeroAttr(attr::kInline)) return false; + + // The function is not able to be inlined if any callee under the CallGraph + // of this function cannot be inlined. + for (const auto& it : *cg_node) { + if (!CanInline(it.second)) { + return false; + } + } + + return true; + } + + // Make a new Relay expression to replace \p expr. + Expr MakeNewExpr(const GlobalVar& global, const Array& args, const Expr& expr) { + ICHECK(expr->IsInstance() || expr->IsInstance()); + auto base_func = call_graph_->GetGlobalFunction(global); + const auto* fn = base_func.as(); + VLOG(1) << "Make new expr " << fn; + ICHECK(fn) << "Expected to work on a Relay function."; + + // There is an inconsistency here, the function itself gets shallow-copied but the body is not + // shallow-copied. + auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, fn->attrs); + // Inline the function body to the caller if this function uses default + // compiler, i.e. no external codegen is needed. + if (!func->GetAttr(attr::kCompiler).defined() && + !func->GetAttr(attr::kExternalSymbol).defined()) { + ICHECK_EQ(func->params.size(), args.size()) + << "Mismatch found in the number of parameters and call args"; + // Bind the parameters with call args. + Map bind_map; + for (size_t i = 0; i < args.size(); i++) { + bind_map.Set(fn->params[i], args[i]); + } + if (const auto* gvn = expr.as()) { + auto ret_type = gvn->checked_type(); + // Cannot replace TensorType/TensorTupleType with FuncType. Therefore, + // we simply inline the function as a closure instead of directly using + // its body when the global var returns FuncType. + return ret_type->IsInstance() ? std::move(func) : func->body; + } else { + ICHECK(expr->IsInstance()); + return Bind(func->body, bind_map); + } + } else if (const auto* call_node = expr.as()) { + return Call(func, args, call_node->attrs, call_node->type_args); + } else { + return std::move(func); + } + } + + /*! + * \brief The current call graph entry that is being handled. Each entry + * contains a global function. + */ + CallGraphEntry* cur_node_; + /*! \brief The call graph that is used for global function lookup. */ + const CallGraphNode* call_graph_; +}; + +IRModule UnmergeComposites(const IRModule& module) { + CallGraph cg(module); + auto topo = cg->TopologicalOrder(); + std::reverse(topo.begin(), topo.end()); + std::unordered_set original_entry; + VLOG_CONTEXT << "Unmerge Composite"; + + VLOG(1) << "Topo size: " << topo.size(); + for (auto* it : topo) { + auto base_func = module->Lookup(it->GetNameHint()); + if (it->GetNameHint() != "main") { + // Check kSymbol that is the correct one + // Check kCompiler that is the correct one + if (const auto* fn = base_func.as()) { + VLOG(1) << "Func name " << it->GetNameHint() << std::endl << "-------" << std::endl; + auto func = GetRef(fn); + auto new_func = Unmerger(it, cg.operator->()).Unmerge(func); + + cg->module->Update(it->GetGlobalVar(), new_func); + } + } + } + VLOG(1) << "Post unmerge module " << std::endl; + VLOG(1) << module; + VLOG(1) << "------------- " << std::endl; + return module; +} + +namespace transform { + +Pass UnmergeComposites() { + runtime::TypedPackedFunc pass_func = + [=](IRModule m, PassContext pc) { return relay::UnmergeComposites(m); }; + return CreateModulePass(pass_func, 1, "UnmergeComposites", {}); +} + +TVM_REGISTER_GLOBAL("relay._transform.UnmergeComposites").set_body_typed(UnmergeComposites); + +} // namespace transform + +} // namespace relay + +} // namespace tvm \ No newline at end of file From 548b5de75c8156ec535245a105129bdba0546729 Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Wed, 23 Mar 2022 08:01:34 -0700 Subject: [PATCH 2/8] Register more patterns for TRT --- python/tvm/relay/op/contrib/tensorrt.py | 306 ++++++++++++++++++++++-- 1 file changed, 284 insertions(+), 22 deletions(-) diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index 986489bae29d..60d596d2ddbe 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -181,7 +181,6 @@ def partition_for_tensorrt( ), transform.FoldConstant(), transform.MergeComposite(pattern_table()), - # transform.AnnotateTarget(["tensorrt"], include_non_call_ops=True), transform.AnnotateTarget("tensorrt"), transform.MergeCompilerRegions(), transform.PartitionGraph(), @@ -490,11 +489,6 @@ def check_add(pattern): # pylint: disable=unused-variable return True -def squeeze_pattern(): - """Create the pattern for squeeze.""" - return is_op("squeeze")(wildcard()) - - def check_squeeze(pattern): attrs, args = pattern.attrs, pattern.args if not is_supported_trt_dtype(args): @@ -513,10 +507,293 @@ def batch_mul_pattern(): return is_op("nn.batch_matmul")(wildcard(), wildcard()) +def check_batch_matmul(pattern): + # attrs, args = pattern.attrs, pattern.args + args = pattern.args + if not is_supported_trt_dtype(args): + return False + if get_tensorrt_use_implicit_batch_mode() and len(pattern.args[0].checked_type.shape) != len( + pattern.args[1].checked_type.shape + ): + logger.info("nn.batch_matmul: requires use_implict_batch=False.") + return False + return True + + +def dense_pattern(): + """Create the pattern for nn.dense.""" + return is_op("nn.dense")(wildcard(), wildcard()) + + +def nn_layer_norm_pattern(): + """Create the pattern for nn.layer_norm.""" + return is_op("nn.layer_norm")(wildcard(), wildcard()) + + +def nn_bias_add_pattern(): + """Create the pattern for nn.bias_add.""" + return is_op("nn.bias_add")(wildcard(), wildcard()) + + +def nn_max_pool2d_pattern(): + """Create the pattern for nn.max_pool2d.""" + return is_op("nn.max_pool2d")(wildcard(), wildcard()) + + +def nn_avg_pool2d_pattern(): + """Create the pattern for nn.avg_pool2d.""" + return is_op("nn.avg_pool2d")(wildcard(), wildcard()) + + +def nn_global_max_pool2d_pattern(): + """Create the pattern for nn.global_max_pool2d.""" + return is_op("nn.global_max_pool2d")(wildcard(), wildcard()) + + +def nn_global_avg_pool2d_pattern(): + """Create the pattern for nn.global_avg_pool2d.""" + return is_op("nn.global_avg_pool2d")(wildcard(), wildcard()) + + +def squeeze_pattern(): + """Create the pattern for squeeze.""" + return is_op("squeeze")(wildcard()) + + +def expand_dims_pattern(): + """Create the pattern for expand_dims.""" + return is_op("expand_dims")(wildcard()) + + +def split_pattern(): + """Create the pattern for split.""" + return is_op("split")(wildcard()) + + +def concatenate_pattern(): + """Create the pattern for concatenate.""" + return is_op("concatenate")(wildcard()) + + +def transpose_pattern(): + """Create the pattern for transpose.""" + return is_op("transpose")(wildcard()) + + +def nn_conv2d_transpose_pattern(): + """Create the pattern for nn.conv2d_transpose.""" + return is_op("nn.conv2d_transpose")(wildcard()) + + +def layout_transform_pattern(): + """Create the pattern for nn.layout_transform.""" + return is_op("layout_transform")(wildcard()) + + +def reshape_pattern(): + """Create the pattern for reshape.""" + return is_op("reshape")(wildcard()) + + +def reshape_pattern(): + """Create the pattern for reshape.""" + return is_op("reshape")(wildcard()) + + +def nn_pad_pattern(): + """Create the pattern for nn.pad.""" + return is_op("nn.pad")(wildcard()) + + +def strided_slice_pattern(): + """Create the pattern for strided_slice.""" + return is_op("strided_slice")(wildcard()) + + +def nn_adaptive_max_pool2d_pattern(): + """Create the pattern for nn.adaptive_max_pool2d.""" + return is_op("nn.adaptive_max_pool2d")(wildcard(), wildcard()) + + +def nn_adaptive_avg_pool2d_pattern(): + """Create the pattern for nn.adaptive_avg_pool2d.""" + return is_op("nn.adaptive_avg_pool2d")(wildcard()) + + +def nn_conv3d_pattern(): + """Create the pattern for nn.conv3d.""" + return is_op("nn.conv3d")(wildcard()) + + +def nn_max_pool3d_pattern(): + """Create the pattern for nn.max_pool3d.""" + return is_op("nn.max_pool3d")(wildcard()) + + +def nn_avg_pool3d_pattern(): + """Create the pattern for nn.avg_pool3d.""" + return is_op("nn.avg_pool3d")(wildcard()) + + +def nn_conv3d_transpose_pattern(): + """Create the pattern for nn.conv3d_transpose.""" + return is_op("nn.conv3d_transpose")(wildcard()) + + +def nn_conv1d_pattern(): + """Create the pattern for nn.conv1d.""" + return is_op("nn.conv1d")(wildcard()) + + +def nn_softmax_pattern(): + """Create the pattern for nn.softmax.""" + return is_op("nn.softmax")(wildcard()) + + +def nn_batch_norm_pattern(): + """Create the pattern for nn.batch_norm.""" + return is_op("nn.batch_norm")(wildcard(), wildcard()) + + +def nn_relu_pattern(): + """Create the pattern for nn.relu.""" + return is_op("nn.relu")(wildcard()) + + +def sigmoid_pattern(): + """Create the pattern for sigmoid.""" + return is_op("sigmoid")(wildcard()) + + +def tanh_pattern(): + """Create the pattern for tanh.""" + return is_op("tanh")(wildcard()) + + +def power_pattern(): + """Create the pattern for power.""" + return is_op("power")(wildcard()) + + +def maximum_pattern(): + """Create the pattern for maximum.""" + return is_op("maximum")(wildcard()) + + +def minimum_pattern(): + """Create the pattern for minimum.""" + return is_op("minimum")(wildcard()) + + +def exp_pattern(): + """Create the pattern for exp.""" + return is_op("exp")(wildcard()) + + +def log_pattern(): + """Create the pattern for log.""" + return is_op("log")(wildcard()) + + +def abs_pattern(): + """Create the pattern for abs.""" + return is_op("abs")(wildcard()) + + +def sqrt_pattern(): + """Create the pattern for sqrt.""" + return is_op("sqrt")(wildcard()) + + +def negative_pattern(): + """Create the pattern for negative.""" + return is_op("negative")(wildcard()) + + +def nn_batch_flatten_pattern(): + """Create the pattern for nn.batch_flatten.""" + return is_op("nn.batch_flatten")(wildcard()) + + +def clip_pattern(): + """Create the pattern for clip.""" + return is_op("clip")(wildcard()) + + +def subtract_pattern(): + """Create the pattern for subtract.""" + return is_op("subtract")(wildcard(), wildcard()) + + +def multiply_pattern(): + """Create the pattern for multiply.""" + return is_op("multiply")(wildcard(), wildcard()) + + +def divide_pattern(): + """Create the pattern for divide.""" + return is_op("divide")(wildcard(), wildcard()) + + +def sum_pattern(): + """Create the pattern for sum.""" + return is_op("sum")(wildcard(), wildcard()) + + +def prod_pattern(): + """Create the pattern for prod.""" + return is_op("prod")(wildcard(), wildcard()) + + +def max_pattern(): + """Create the pattern for max.""" + return is_op("max")(wildcard(), wildcard()) + + +def mix_pattern(): + """Create the pattern for min.""" + return is_op("min")(wildcard(), wildcard()) + + +def clip_pattern(): + """Create the pattern for clip.""" + return is_op("clip")(wildcard()) + + +def sin_pattern(): + """Create the pattern for sin.""" + return is_op("sin")(wildcard()) + + +def cos_pattern(): + """Create the pattern for cos.""" + return is_op("cos")(wildcard()) + + +def atan_pattern(): + """Create the pattern for atan.""" + return is_op("atan")(wildcard()) + + +def ceil_pattern(): + """Create the pattern for ceil.""" + return is_op("ceil")(wildcard()) + + +def erf_pattern(): + """Create the pattern for erf.""" + return is_op("erf")(wildcard()) + + +def nn_leaky_relu_pattern(): + """Create the pattern for nn.leaky_relu.""" + return is_op("nn.leaky_relu")(wildcard()) + + @register_pattern_table("tensorrt") def pattern_table(): return [("tensorrt.nn.conv2d", conv2d_pattern(), check_conv2d), ("tensorrt.squeeze", squeeze_pattern(), check_squeeze), - ("tensorrt.add", add_pattern(), check_add), ("tensorrt.nn.batch_matmul", batch_mul_pattern())] + ("tensorrt.add", add_pattern(), check_add), ("tensorrt.nn.batch_matmul", batch_mul_pattern(), check_batch_matmul), ("tensorrt.nn.dense", dense_pattern())] @_register_external_dynamic_check_func("nn.dense") @@ -538,21 +815,6 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable return True -@_register_external_dynamic_check_func("tensorrt.nn.batch_matmul") -def batch_matmul_annotate_fn(expr): - """Check if dense is supported by TensorRT.""" - - args = expr.args - if not is_supported_trt_dtype(args): - return False - if get_tensorrt_use_implicit_batch_mode() and len(expr.args[0].checked_type.shape) != len( - expr.args[1].checked_type.shape - ): - logger.info("nn.batch_matmul: requires use_implict_batch=False.") - return False - return True - - @_register_external_dynamic_check_func("nn.layer_norm") def layer_norm_annotate_fn(expr): """Check if dense is supported by TensorRT.""" From 4b1022553d7a8f830f846d017de48ad20695b89c Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Wed, 23 Mar 2022 09:07:00 -0700 Subject: [PATCH 3/8] Add more patterns --- python/tvm/relay/op/contrib/tensorrt.py | 313 ++++-------------------- 1 file changed, 45 insertions(+), 268 deletions(-) diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index 60d596d2ddbe..806b1987e900 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -520,280 +520,57 @@ def check_batch_matmul(pattern): return True -def dense_pattern(): - """Create the pattern for nn.dense.""" - return is_op("nn.dense")(wildcard(), wildcard()) - - -def nn_layer_norm_pattern(): - """Create the pattern for nn.layer_norm.""" - return is_op("nn.layer_norm")(wildcard(), wildcard()) - - -def nn_bias_add_pattern(): - """Create the pattern for nn.bias_add.""" - return is_op("nn.bias_add")(wildcard(), wildcard()) - - -def nn_max_pool2d_pattern(): - """Create the pattern for nn.max_pool2d.""" - return is_op("nn.max_pool2d")(wildcard(), wildcard()) - - -def nn_avg_pool2d_pattern(): - """Create the pattern for nn.avg_pool2d.""" - return is_op("nn.avg_pool2d")(wildcard(), wildcard()) - - -def nn_global_max_pool2d_pattern(): - """Create the pattern for nn.global_max_pool2d.""" - return is_op("nn.global_max_pool2d")(wildcard(), wildcard()) - - -def nn_global_avg_pool2d_pattern(): - """Create the pattern for nn.global_avg_pool2d.""" - return is_op("nn.global_avg_pool2d")(wildcard(), wildcard()) - - -def squeeze_pattern(): - """Create the pattern for squeeze.""" - return is_op("squeeze")(wildcard()) - - -def expand_dims_pattern(): - """Create the pattern for expand_dims.""" - return is_op("expand_dims")(wildcard()) - - -def split_pattern(): - """Create the pattern for split.""" - return is_op("split")(wildcard()) - - -def concatenate_pattern(): - """Create the pattern for concatenate.""" - return is_op("concatenate")(wildcard()) - - -def transpose_pattern(): - """Create the pattern for transpose.""" - return is_op("transpose")(wildcard()) - - -def nn_conv2d_transpose_pattern(): - """Create the pattern for nn.conv2d_transpose.""" - return is_op("nn.conv2d_transpose")(wildcard()) - - -def layout_transform_pattern(): - """Create the pattern for nn.layout_transform.""" - return is_op("layout_transform")(wildcard()) - - -def reshape_pattern(): - """Create the pattern for reshape.""" - return is_op("reshape")(wildcard()) - - -def reshape_pattern(): - """Create the pattern for reshape.""" - return is_op("reshape")(wildcard()) - - -def nn_pad_pattern(): - """Create the pattern for nn.pad.""" - return is_op("nn.pad")(wildcard()) - - -def strided_slice_pattern(): - """Create the pattern for strided_slice.""" - return is_op("strided_slice")(wildcard()) - - -def nn_adaptive_max_pool2d_pattern(): - """Create the pattern for nn.adaptive_max_pool2d.""" - return is_op("nn.adaptive_max_pool2d")(wildcard(), wildcard()) - - -def nn_adaptive_avg_pool2d_pattern(): - """Create the pattern for nn.adaptive_avg_pool2d.""" - return is_op("nn.adaptive_avg_pool2d")(wildcard()) - - -def nn_conv3d_pattern(): - """Create the pattern for nn.conv3d.""" - return is_op("nn.conv3d")(wildcard()) - - -def nn_max_pool3d_pattern(): - """Create the pattern for nn.max_pool3d.""" - return is_op("nn.max_pool3d")(wildcard()) - - -def nn_avg_pool3d_pattern(): - """Create the pattern for nn.avg_pool3d.""" - return is_op("nn.avg_pool3d")(wildcard()) - - -def nn_conv3d_transpose_pattern(): - """Create the pattern for nn.conv3d_transpose.""" - return is_op("nn.conv3d_transpose")(wildcard()) - - -def nn_conv1d_pattern(): - """Create the pattern for nn.conv1d.""" - return is_op("nn.conv1d")(wildcard()) - - -def nn_softmax_pattern(): - """Create the pattern for nn.softmax.""" - return is_op("nn.softmax")(wildcard()) - - -def nn_batch_norm_pattern(): - """Create the pattern for nn.batch_norm.""" - return is_op("nn.batch_norm")(wildcard(), wildcard()) - - -def nn_relu_pattern(): - """Create the pattern for nn.relu.""" - return is_op("nn.relu")(wildcard()) - - -def sigmoid_pattern(): - """Create the pattern for sigmoid.""" - return is_op("sigmoid")(wildcard()) - - -def tanh_pattern(): - """Create the pattern for tanh.""" - return is_op("tanh")(wildcard()) - - -def power_pattern(): - """Create the pattern for power.""" - return is_op("power")(wildcard()) - - -def maximum_pattern(): - """Create the pattern for maximum.""" - return is_op("maximum")(wildcard()) - - -def minimum_pattern(): - """Create the pattern for minimum.""" - return is_op("minimum")(wildcard()) - - -def exp_pattern(): - """Create the pattern for exp.""" - return is_op("exp")(wildcard()) - - -def log_pattern(): - """Create the pattern for log.""" - return is_op("log")(wildcard()) - - -def abs_pattern(): - """Create the pattern for abs.""" - return is_op("abs")(wildcard()) - - -def sqrt_pattern(): - """Create the pattern for sqrt.""" - return is_op("sqrt")(wildcard()) - - -def negative_pattern(): - """Create the pattern for negative.""" - return is_op("negative")(wildcard()) - - -def nn_batch_flatten_pattern(): - """Create the pattern for nn.batch_flatten.""" - return is_op("nn.batch_flatten")(wildcard()) - - -def clip_pattern(): - """Create the pattern for clip.""" - return is_op("clip")(wildcard()) - - -def subtract_pattern(): - """Create the pattern for subtract.""" - return is_op("subtract")(wildcard(), wildcard()) - - -def multiply_pattern(): - """Create the pattern for multiply.""" - return is_op("multiply")(wildcard(), wildcard()) - - -def divide_pattern(): - """Create the pattern for divide.""" - return is_op("divide")(wildcard(), wildcard()) - - -def sum_pattern(): - """Create the pattern for sum.""" - return is_op("sum")(wildcard(), wildcard()) - - -def prod_pattern(): - """Create the pattern for prod.""" - return is_op("prod")(wildcard(), wildcard()) - - -def max_pattern(): - """Create the pattern for max.""" - return is_op("max")(wildcard(), wildcard()) - - -def mix_pattern(): - """Create the pattern for min.""" - return is_op("min")(wildcard(), wildcard()) - - -def clip_pattern(): - """Create the pattern for clip.""" - return is_op("clip")(wildcard()) - - -def sin_pattern(): - """Create the pattern for sin.""" - return is_op("sin")(wildcard()) - - -def cos_pattern(): - """Create the pattern for cos.""" - return is_op("cos")(wildcard()) - - -def atan_pattern(): - """Create the pattern for atan.""" - return is_op("atan")(wildcard()) - - -def ceil_pattern(): - """Create the pattern for ceil.""" - return is_op("ceil")(wildcard()) - - -def erf_pattern(): - """Create the pattern for erf.""" - return is_op("erf")(wildcard()) +def unary_op_pattern(op): + """Matches unary operation""" + pattern = is_op(op)( + wildcard() + ) + return pattern -def nn_leaky_relu_pattern(): - """Create the pattern for nn.leaky_relu.""" - return is_op("nn.leaky_relu")(wildcard()) +def binary_op_pattern(op): + """Matches binary operation""" + pattern = is_op(op)( + wildcard(), + wildcard() + ) + return pattern @register_pattern_table("tensorrt") def pattern_table(): - return [("tensorrt.nn.conv2d", conv2d_pattern(), check_conv2d), ("tensorrt.squeeze", squeeze_pattern(), check_squeeze), - ("tensorrt.add", add_pattern(), check_add), ("tensorrt.nn.batch_matmul", batch_mul_pattern(), check_batch_matmul), ("tensorrt.nn.dense", dense_pattern())] + return [ + ("tensorrt.nn.conv2d", conv2d_pattern(), check_conv2d), + ("tensorrt.squeeze", binary_op_pattern("squeeze"), check_squeeze), + ("tensorrt.add", binary_op_pattern("add"), check_add), + ("tensorrt.nn.dense", unary_op_pattern("nn.dense")), + ("tensorrt.bias_add", binary_op_pattern("nn.bias_add")), + ("tensorrt.nn.batch_matmul", binary_op_pattern("nn.batch_matmul")), + ("tensorrt.substract", binary_op_pattern("substract")), + ("tensorrt.divide", binary_op_pattern("divide")), + ("tensorrt.multiply", binary_op_pattern("multiply")), + ("tensorrt.split", unary_op_pattern("split")), + ("tensorrt.reshape", unary_op_pattern("reshape")), + ("tensorrt.nn.relu", unary_op_pattern("nn.relu")), + ("tensorrt.nn.leaky.relu", unary_op_pattern("nn.leaky.relu")), + ("tensorrt.nn.pad", unary_op_pattern("nn.pad")), + ("tensorrt.sigmoid", unary_op_pattern("sigmoid")), + ("tensorrt.tanh", unary_op_pattern("tanh")), + ("tensorrt.exp", unary_op_pattern("exp")), + ("tensorrt.log", unary_op_pattern("log")), + ("tensorrt.sqrt", unary_op_pattern("sqrt")), + ("tensorrt.abs", unary_op_pattern("abs")), + ("tensorrt.negative", unary_op_pattern("negative")), + ("tensorrt.sin", unary_op_pattern("sin")), + ("tensorrt.cos", unary_op_pattern("cos")), + ("tensorrt.atan", unary_op_pattern("atan")), + ("tensorrt.ceil", unary_op_pattern("ceil")), + ("tensorrt.floor", unary_op_pattern("floor")), + ("tensorrt.erf", unary_op_pattern("erf")), + ("tensorrt.nn.softmax", unary_op_pattern("nn.softmax")), + ("tensorrt.nn.layer_norm", unary_op_pattern("nn.layer_norm")), + ("tensorrt.nn.max_pool2d", unary_op_pattern("nn.max_pool2d")), + ] @_register_external_dynamic_check_func("nn.dense") From 067fc878bad5d1b9b5f7173ee9e47eeb493d1bbc Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Wed, 23 Mar 2022 10:43:49 -0700 Subject: [PATCH 4/8] [WIP] Pre-cleanup on checks --- python/tvm/relay/op/contrib/tensorrt.py | 17 +++-------------- 1 file changed, 3 insertions(+), 14 deletions(-) diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index 806b1987e900..e8a99d91e8c4 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -430,12 +430,8 @@ def conv1d_annotate_fn(expr): # pylint: disable=unused-variable return True -def conv2d_pattern(): - conv2d = is_op("nn.conv2d")(wildcard(), wildcard()) - return conv2d - - -def check_conv2d(pattern): +@_register_external_dynamic_check_func("nn.conv2d") +def conv2d_annotate_fn(pattern): attrs, args = pattern.attrs, pattern.args if not is_supported_trt_dtype(args): return False @@ -502,11 +498,6 @@ def check_squeeze(pattern): return True -def batch_mul_pattern(): - """Create the pattern for squeeze.""" - return is_op("nn.batch_matmul")(wildcard(), wildcard()) - - def check_batch_matmul(pattern): # attrs, args = pattern.attrs, pattern.args args = pattern.args @@ -540,19 +531,17 @@ def binary_op_pattern(op): @register_pattern_table("tensorrt") def pattern_table(): return [ - ("tensorrt.nn.conv2d", conv2d_pattern(), check_conv2d), + ("tensorrt.nn.conv2d", binary_op_pattern("nn.conv2d")), ("tensorrt.squeeze", binary_op_pattern("squeeze"), check_squeeze), ("tensorrt.add", binary_op_pattern("add"), check_add), ("tensorrt.nn.dense", unary_op_pattern("nn.dense")), ("tensorrt.bias_add", binary_op_pattern("nn.bias_add")), ("tensorrt.nn.batch_matmul", binary_op_pattern("nn.batch_matmul")), - ("tensorrt.substract", binary_op_pattern("substract")), ("tensorrt.divide", binary_op_pattern("divide")), ("tensorrt.multiply", binary_op_pattern("multiply")), ("tensorrt.split", unary_op_pattern("split")), ("tensorrt.reshape", unary_op_pattern("reshape")), ("tensorrt.nn.relu", unary_op_pattern("nn.relu")), - ("tensorrt.nn.leaky.relu", unary_op_pattern("nn.leaky.relu")), ("tensorrt.nn.pad", unary_op_pattern("nn.pad")), ("tensorrt.sigmoid", unary_op_pattern("sigmoid")), ("tensorrt.tanh", unary_op_pattern("tanh")), From a4d95c7d949441c887525c9d6017bd84883fa7ce Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Thu, 24 Mar 2022 04:51:23 -0700 Subject: [PATCH 5/8] More cleanup and op support added for TRT patterns --- python/tvm/relay/op/contrib/tensorrt.py | 376 ++++++++++----------- python/tvm/relay/transform/transform.py | 20 +- src/relay/transforms/inline_composites.cc | 119 +++++++ src/relay/transforms/unmerge_composites.cc | 254 -------------- 4 files changed, 303 insertions(+), 466 deletions(-) create mode 100644 src/relay/transforms/inline_composites.cc delete mode 100644 src/relay/transforms/unmerge_composites.cc diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index e8a99d91e8c4..eb814ab5012a 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -28,6 +28,7 @@ from tvm.relay.expr import Call, Constant, GlobalVar, Tuple, TupleGetItem, Var from tvm.relay.expr_functor import ExprMutator, ExprVisitor from tvm.relay.op.contrib.register import register_pattern_table +from tvm.relay.op.transform import split logger = logging.getLogger("TensorRT") supported_types = ["float32", "float16"] @@ -41,8 +42,7 @@ def is_supported_trt_dtype(args): True if supported, False if not. """ if not all([x.checked_type.dtype in supported_types for x in args]): - logger.info( - "Only float32 and float16 inputs are supported for TensorRT BYOC.") + logger.info("Only float32 and float16 inputs are supported for TensorRT BYOC.") return False return True @@ -54,8 +54,7 @@ def is_tensorrt_runtime_enabled(): ret: bool True if present, False if not. """ - check_enabled = tvm.get_global_func( - "relay.op.is_tensorrt_runtime_enabled", True) + check_enabled = tvm.get_global_func("relay.op.is_tensorrt_runtime_enabled", True) if check_enabled: return check_enabled() return False @@ -154,8 +153,7 @@ def partition_for_tensorrt( assert isinstance(version, tuple) and len(version) == 3 config["tensorrt_version"] = version else: - linked_version = tuple(tvm.get_global_func( - "relay.op.get_tensorrt_version")()) + linked_version = tuple(tvm.get_global_func("relay.op.get_tensorrt_version")()) if not linked_version: logger.warning( "TVM was not built against TensorRT and no version was provided to " @@ -184,13 +182,12 @@ def partition_for_tensorrt( transform.AnnotateTarget("tensorrt"), transform.MergeCompilerRegions(), transform.PartitionGraph(), - transform.UnmergeComposites(), + transform.InlineComposites("tensorrt"), transform.InferType(), ] ) with tvm.transform.PassContext(opt_level=3, config={"relay.ext.tensorrt.options": config}): mod = seq(mod) - print(mod) mod = prune_tensorrt_subgraphs(mod) return mod, config @@ -219,14 +216,13 @@ def check_dynamism(args, op_name): elif isinstance(arg, Tuple): return check_dynamism(arg.fields, op_name) else: - logger.info( - "Arg not supported in TensorRT for %s with type %s", op_name, type(arg)) + logger.info("Arg not supported in TensorRT for %s with type %s", op_name, type(arg)) return True return False def _register_external_op_helper_with_checker(op_name, checker): - @ tvm.ir.register_op_attr(op_name, "target.tensorrt") + @tvm.ir.register_op_attr(op_name, "target.tensorrt") def _func_wrapper(expr): attrs, args = expr.attrs, expr.args # ops with dynamic shapes are offloaded to VM @@ -246,8 +242,7 @@ def _func_wrapper(expr): # have been excluded because they occur in PT MaskRCNN model. The long term solution is # to switch to explicit batch mode after performance regressions are solved. if all( - [list(map(int, shape)) in [[300, 64, 7, 7], [300, 1, 1, 1]] - for shape in shapes] + [list(map(int, shape)) in [[300, 64, 7, 7], [300, 1, 1, 1]] for shape in shapes] ): return False return checker(attrs, args, op_name) @@ -324,8 +319,7 @@ def trt_version_annotate_fn(version): def _func_wrapper(attrs, args, op_name): if get_tensorrt_version() < version: logger.info( - "%s: requires TensorRT version %s or higher.", op_name, ".".join( - map(str, version)) + "%s: requires TensorRT version %s or higher.", op_name, ".".join(map(str, version)) ) return False return True @@ -333,18 +327,12 @@ def _func_wrapper(attrs, args, op_name): return _func_wrapper -_register_external_op_helper_with_checker( - "nn.leaky_relu", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker( - "sin", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker( - "cos", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker( - "atan", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker( - "ceil", trt_version_annotate_fn((5, 1, 5))) -_register_external_op_helper_with_checker( - "erf", trt_version_annotate_fn((7, 0, 0))) +_register_external_op_helper_with_checker("nn.leaky_relu", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker("sin", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker("cos", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker("atan", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker("ceil", trt_version_annotate_fn((5, 1, 5))) +_register_external_op_helper_with_checker("erf", trt_version_annotate_fn((7, 0, 0))) @_register_external_dynamic_check_func("add") @@ -355,8 +343,7 @@ def add_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False shapes = [ - [int(x) if not isinstance(x, tvm.tir.expr.Any) - else -1 for x in arg.checked_type.shape] + [int(x) if not isinstance(x, tvm.tir.expr.Any) else -1 for x in arg.checked_type.shape] for arg in args ] @@ -386,15 +373,13 @@ def batch_norm_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if len(args[0].checked_type.shape) == 5 and get_tensorrt_version() < (6, 0, 1): - logger.info( - "nn.batch_norm: TensorRT 6.0.1 or higher is required for rank 5 inputs.") + logger.info("nn.batch_norm: TensorRT 6.0.1 or higher is required for rank 5 inputs.") return False if len(args[0].checked_type.shape) > 5: logger.info("nn.batch_norm: Input rank must be 5 or less.") return False if int(attrs.axis) not in (1, 3): - logger.info("nn.batch_norm: axis is %d but must be 1 or 3.", - int(attrs.axis)) + logger.info("nn.batch_norm: axis is %d but must be 1 or 3.", int(attrs.axis)) return False return True @@ -420,148 +405,33 @@ def conv1d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCW": - logger.info("nn.conv1d: data_layout is %s but must be NCW.", - attrs.data_layout) + logger.info("nn.conv1d: data_layout is %s but must be NCW.", attrs.data_layout) return False if attrs.kernel_layout != "OIW": - logger.info("nn.conv1d: kernel_layout is %s but must be OIW.", - attrs.kernel_layout) + logger.info("nn.conv1d: kernel_layout is %s but must be OIW.", attrs.kernel_layout) return False return True @_register_external_dynamic_check_func("nn.conv2d") -def conv2d_annotate_fn(pattern): - attrs, args = pattern.attrs, pattern.args +def conv2d_annotate_fn(expr): # pylint: disable=unused-variable + """Check if nn.conv2d is supported by TensorRT.""" + + attrs, args = expr.attrs, expr.args if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCHW": - logger.info( - "nn.conv2d: data_layout is %s but must be NCHW.", attrs.data_layout) + logger.info("nn.conv2d: data_layout is %s but must be NCHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIHW": - logger.info( - "nn.conv2d: kernel_layout is %s but must be OIHW.", attrs.kernel_layout) + logger.info("nn.conv2d: kernel_layout is %s but must be OIHW.", attrs.kernel_layout) return False if attrs.out_layout and attrs.out_layout != "NCHW": - logger.info( - "nn.conv2d: out_layout is %s but must be NCHW.", attrs.out_layout) - return False - return True - - -def add_pattern(): - pat = is_op("add")(wildcard(), wildcard()) - return pat - - -def check_add(pattern): # pylint: disable=unused-variable - """Check if add is supported by TensorRT.""" - - args = pattern.args - if not is_supported_trt_dtype(args): - return False - shapes = [ - [int(x) if not isinstance(x, tvm.tir.expr.Any) - else -1 for x in arg.checked_type.shape] - for arg in args - ] - - # Scalars require explicit batch mode. - if get_tensorrt_use_implicit_batch_mode() and any([len(shape) < 1 for shape in shapes]): - return False - - if ( - not get_tensorrt_use_implicit_batch_mode() - and (isinstance(args[0], Constant) or isinstance(args[1], Constant)) - and len(shapes[0]) > 0 - and len(shapes[1]) > 0 - and shapes[0][0] == shapes[1][0] - and shapes[0][0] != 1 - and (len(shapes[0]) > 3 or len(shapes[1]) > 3) - ): - logger.info("add: bug in TRT with adding batched constants.") - return False - return True - - -def check_squeeze(pattern): - attrs, args = pattern.attrs, pattern.args - if not is_supported_trt_dtype(args): - return False - if not attrs.axis: - logger.info("squeeze: must explicitly set axis.") - return False - if get_tensorrt_use_implicit_batch_mode() and any([axis == 0 for axis in map(int, attrs.axis)]): - logger.info("squeeze: can't modify batch dimension.") + logger.info("nn.conv2d: out_layout is %s but must be NCHW.", attrs.out_layout) return False return True -def check_batch_matmul(pattern): - # attrs, args = pattern.attrs, pattern.args - args = pattern.args - if not is_supported_trt_dtype(args): - return False - if get_tensorrt_use_implicit_batch_mode() and len(pattern.args[0].checked_type.shape) != len( - pattern.args[1].checked_type.shape - ): - logger.info("nn.batch_matmul: requires use_implict_batch=False.") - return False - return True - - -def unary_op_pattern(op): - """Matches unary operation""" - pattern = is_op(op)( - wildcard() - ) - return pattern - - -def binary_op_pattern(op): - """Matches binary operation""" - pattern = is_op(op)( - wildcard(), - wildcard() - ) - return pattern - - -@register_pattern_table("tensorrt") -def pattern_table(): - return [ - ("tensorrt.nn.conv2d", binary_op_pattern("nn.conv2d")), - ("tensorrt.squeeze", binary_op_pattern("squeeze"), check_squeeze), - ("tensorrt.add", binary_op_pattern("add"), check_add), - ("tensorrt.nn.dense", unary_op_pattern("nn.dense")), - ("tensorrt.bias_add", binary_op_pattern("nn.bias_add")), - ("tensorrt.nn.batch_matmul", binary_op_pattern("nn.batch_matmul")), - ("tensorrt.divide", binary_op_pattern("divide")), - ("tensorrt.multiply", binary_op_pattern("multiply")), - ("tensorrt.split", unary_op_pattern("split")), - ("tensorrt.reshape", unary_op_pattern("reshape")), - ("tensorrt.nn.relu", unary_op_pattern("nn.relu")), - ("tensorrt.nn.pad", unary_op_pattern("nn.pad")), - ("tensorrt.sigmoid", unary_op_pattern("sigmoid")), - ("tensorrt.tanh", unary_op_pattern("tanh")), - ("tensorrt.exp", unary_op_pattern("exp")), - ("tensorrt.log", unary_op_pattern("log")), - ("tensorrt.sqrt", unary_op_pattern("sqrt")), - ("tensorrt.abs", unary_op_pattern("abs")), - ("tensorrt.negative", unary_op_pattern("negative")), - ("tensorrt.sin", unary_op_pattern("sin")), - ("tensorrt.cos", unary_op_pattern("cos")), - ("tensorrt.atan", unary_op_pattern("atan")), - ("tensorrt.ceil", unary_op_pattern("ceil")), - ("tensorrt.floor", unary_op_pattern("floor")), - ("tensorrt.erf", unary_op_pattern("erf")), - ("tensorrt.nn.softmax", unary_op_pattern("nn.softmax")), - ("tensorrt.nn.layer_norm", unary_op_pattern("nn.layer_norm")), - ("tensorrt.nn.max_pool2d", unary_op_pattern("nn.max_pool2d")), - ] - - @_register_external_dynamic_check_func("nn.dense") def dense_annotate_fn(expr): # pylint: disable=unused-variable """Check if dense is supported by TensorRT.""" @@ -572,8 +442,7 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable input_rank = len(args[0].checked_type.shape) weight_rank = len(args[1].checked_type.shape) if input_rank not in (2, 3, 4): - logger.info( - "nn.dense: input has rank %d but must be 2, 3 or 4.", input_rank) + logger.info("nn.dense: input has rank %d but must be 2, 3 or 4.", input_rank) return False if weight_rank != 2: logger.info("nn.dense: weight has rank %d but must be 2.", weight_rank) @@ -581,6 +450,21 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable return True +@_register_external_dynamic_check_func("nn.batch_matmul") +def batch_matmul_annotate_fn(expr): # pylint: disable=unused-variable + """Check if dense is supported by TensorRT.""" + + args = expr.args + if not is_supported_trt_dtype(args): + return False + if get_tensorrt_use_implicit_batch_mode() and len(expr.args[0].checked_type.shape) != len( + expr.args[1].checked_type.shape + ): + logger.info("nn.batch_matmul: requires use_implict_batch=False.") + return False + return True + + @_register_external_dynamic_check_func("nn.layer_norm") def layer_norm_annotate_fn(expr): """Check if dense is supported by TensorRT.""" @@ -603,8 +487,7 @@ def bias_add_annotate_fn(expr): # pylint: disable=unused-variable return False input_rank = len(args[0].checked_type.shape) if input_rank not in (2, 3, 4): - logger.info( - "nn.bias_add: input rank is %d but must be 2, 3 or 4.", input_rank) + logger.info("nn.bias_add: input rank is %d but must be 2, 3 or 4.", input_rank) return False return True @@ -617,12 +500,10 @@ def max_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info( - "nn.max_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info("nn.max_pool2d: layout is %s but must be NCHW.", attrs.layout) return False if attrs.ceil_mode and get_tensorrt_version() < (5, 1, 5): - logger.info( - "nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") + logger.info("nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") return False return True @@ -635,8 +516,7 @@ def avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info( - "nn.avg_pool2d: layout is %d but must be NCHW.", attrs.layout) + logger.info("nn.avg_pool2d: layout is %d but must be NCHW.", attrs.layout) return False if ( attrs.count_include_pad @@ -652,8 +532,7 @@ def avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.ceil_mode and get_tensorrt_version() < (5, 1, 5): - logger.info( - "nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") + logger.info("nn.avg_pool2d: ceil_mode=True requires TensorRT 5.1.5 or greater.") return False return True @@ -666,8 +545,7 @@ def global_max_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info( - "nn.global_max_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info("nn.global_max_pool2d: layout is %s but must be NCHW.", attrs.layout) return False return True @@ -680,8 +558,7 @@ def global_avg_pool_2d_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.layout != "NCHW": - logger.info( - "nn.global_avg_pool2d: layout is %s but must be NCHW.", attrs.layout) + logger.info("nn.global_avg_pool2d: layout is %s but must be NCHW.", attrs.layout) return False return True @@ -721,8 +598,7 @@ def concatenate_annotate_fn(expr): # pylint: disable=unused-variable attrs, args = expr.attrs, expr.args if any([x.dtype not in supported_types for x in args[0].checked_type.fields]): - logger.info( - "Only float16 and float32 inputs are supported for TensorRT.") + logger.info("Only float16 and float32 inputs are supported for TensorRT.") if not get_tensorrt_use_implicit_batch_mode(): return True if int(attrs.axis) == 0: @@ -731,8 +607,7 @@ def concatenate_annotate_fn(expr): # pylint: disable=unused-variable if isinstance(args[0], Tuple): for tuple_input in args[0].fields: if isinstance(tuple_input, Constant): - logger.info( - "concatenate: can't concatenate tensors with constants.") + logger.info("concatenate: can't concatenate tensors with constants.") return False return True @@ -758,8 +633,7 @@ def conv2d_transpose_annotate_fn(expr): # pylint: disable=unused-variable if not is_supported_trt_dtype(args): return False if attrs.data_layout != "NCHW": - logger.info( - "nn.conv2d_transpose: data_layout is %s but must be NCHW.", attrs.data_layout) + logger.info("nn.conv2d_transpose: data_layout is %s but must be NCHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIHW": logger.info( @@ -767,8 +641,7 @@ def conv2d_transpose_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.out_layout and attrs.out_layout != "NCHW": - logger.info( - "nn.conv2d_transpose: out_layout is %s but must be NCHW.", attrs.out_layout) + logger.info("nn.conv2d_transpose: out_layout is %s but must be NCHW.", attrs.out_layout) return False if attrs.dilation and any([rate != 1 for rate in map(int, attrs.dilation)]): logger.info("nn.conv2d_transpose: dilation rate must be 1.") @@ -857,8 +730,7 @@ def reshape_annotate_fn(expr): # pylint: disable=unused-variable # Resolve -1. for i, value in enumerate(new_shape): if value == -1: - new_shape[i] = original_volume // np.prod( - [x for x in new_shape if x != -1]) + new_shape[i] = original_volume // np.prod([x for x in new_shape if x != -1]) # Remove batch dimension and see if volumes match if shape[0] != new_shape[0]: logger.info("reshape: can't modify batch dimension.") @@ -877,8 +749,7 @@ def pad_annotate_fn(expr): # pylint: disable=unused-variable assert isinstance(pad_value, relay.Constant) pad_value = pad_value.data.numpy().item() if attrs.pad_mode != "constant": - logger.info("nn.pad: pad mode is %s but must be constant.", - attrs.pad_mode) + logger.info("nn.pad: pad mode is %s but must be constant.", attrs.pad_mode) return False if pad_value > 0.0: logger.info("nn.pad: pad value is %f but must be 0.0.", pad_value) @@ -905,8 +776,7 @@ def strided_slice_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((5, 1, 5))(attrs, args, "strided_slice"): return False if get_tensorrt_use_implicit_batch_mode(): - batch_dim_begin_modified = attrs.begin[0] is not None and int( - attrs.begin[0]) != 0 + batch_dim_begin_modified = attrs.begin[0] is not None and int(attrs.begin[0]) != 0 batch_dim_end_modified = ( attrs.end[0] is not None and int(attrs.end[0]) != -1 @@ -979,16 +849,13 @@ def conv3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.conv3d"): return False if attrs.data_layout != "NCDHW": - logger.info( - "nn.conv3d: data_layout is %s but must be NCDHW.", attrs.data_layout) + logger.info("nn.conv3d: data_layout is %s but must be NCDHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIDHW": - logger.info( - "nn.conv3d: kernel_layout is %s but must be OIDHW.", attrs.kernel_layout) + logger.info("nn.conv3d: kernel_layout is %s but must be OIDHW.", attrs.kernel_layout) return False if attrs.out_layout and attrs.out_layout != "NCDHW": - logger.info( - "nn.conv3d: out_layout is %s but must be NCDHW.", attrs.out_layout) + logger.info("nn.conv3d: out_layout is %s but must be NCDHW.", attrs.out_layout) return False return True @@ -1003,8 +870,7 @@ def max_pool_3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.max_pool3d"): return False if attrs.layout != "NCDHW": - logger.info( - "nn.max_pool3d: layout is %s but must be NCDHW.", attrs.layout) + logger.info("nn.max_pool3d: layout is %s but must be NCDHW.", attrs.layout) return False return True @@ -1019,8 +885,7 @@ def avg_pool_3d_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.avg_pool3d"): return False if attrs.layout != "NCDHW": - logger.info( - "nn.avg_pool3d: layout is %s but must be NCDHW.", attrs.layout) + logger.info("nn.avg_pool3d: layout is %s but must be NCDHW.", attrs.layout) return False return True @@ -1035,8 +900,7 @@ def conv3d_transpose_annotate_fn(expr): # pylint: disable=unused-variable if not trt_version_annotate_fn((6, 0, 1))(attrs, args, "nn.conv3d_transpose"): return False if attrs.data_layout != "NCDHW": - logger.info( - "nn.conv3d_transpose: data_layout is %s but must be NCDHW.", attrs.data_layout) + logger.info("nn.conv3d_transpose: data_layout is %s but must be NCDHW.", attrs.data_layout) return False if attrs.kernel_layout != "OIDHW": logger.info( @@ -1044,8 +908,7 @@ def conv3d_transpose_annotate_fn(expr): # pylint: disable=unused-variable ) return False if attrs.out_layout and attrs.out_layout != "NCDHW": - logger.info( - "nn.conv3d_transpose: out_layout is %s but must be NCDHW.", attrs.out_layout) + logger.info("nn.conv3d_transpose: out_layout is %s but must be NCDHW.", attrs.out_layout) return False if attrs.dilation and any([rate != 1 for rate in map(int, attrs.dilation)]): logger.info("nn.conv3d_transpose: dilation rate must be 1.") @@ -1056,6 +919,114 @@ def conv3d_transpose_annotate_fn(expr): # pylint: disable=unused-variable return True +def unary_op_pattern(op): + """Matches unary operation""" + pattern = is_op(op)(wildcard()) + return pattern + + +def binary_op_pattern(op): + """Matches binary operation""" + pattern = is_op(op)(wildcard(), wildcard()) + return pattern + + +@register_pattern_table("tensorrt") +def pattern_table(): + """Get the Tensorrt compiler pattern table for supported ops.""" + + return [ + ("tensorrt.nn.conv3d", binary_op_pattern("nn.conv3d"), conv3d_annotate_fn), + ("tensorrt.nn.conv2d", binary_op_pattern("nn.conv2d"), conv2d_annotate_fn), + ("tensorrt.nn.conv1d", binary_op_pattern("nn.conv1d"), conv1d_annotate_fn), + ( + "tensorrt.nn.conv2d_transpose", + binary_op_pattern("nn.conv2d_transpose"), + conv2d_transpose_annotate_fn, + ), + ("tensorrt.squeeze", binary_op_pattern("squeeze"), squeeze_annotate_fn), + ("tensorrt.add", binary_op_pattern("add"), add_annotate_fn), + ("tensorrt.nn.dense", unary_op_pattern("nn.dense"), dense_annotate_fn), + ("tensorrt.bias_add", binary_op_pattern("nn.bias_add"), bias_add_annotate_fn), + ( + "tensorrt.nn.batch_matmul", + binary_op_pattern("nn.batch_matmul"), + batch_matmul_annotate_fn, + ), + ("tensorrt.divide", binary_op_pattern("divide")), + ("tensorrt.multiply", binary_op_pattern("multiply")), + ("tensorrt.split", unary_op_pattern("split")), + ("tensorrt.reshape", unary_op_pattern("reshape")), + ("tensorrt.nn.relu", unary_op_pattern("nn.relu")), + ( + "tensorrt.nn.leaky_relu", + unary_op_pattern("nn.leaky_relu"), + trt_version_annotate_fn((5, 1, 5)), + ), + ("tensorrt.nn.pad", unary_op_pattern("nn.pad")), + ("tensorrt.sigmoid", unary_op_pattern("sigmoid")), + ("tensorrt.tanh", unary_op_pattern("tanh")), + ("tensorrt.exp", unary_op_pattern("exp")), + ("tensorrt.log", unary_op_pattern("log")), + ("tensorrt.sqrt", unary_op_pattern("sqrt")), + ("tensorrt.abs", unary_op_pattern("abs")), + ("tensorrt.power", unary_op_pattern("power")), + ("tensorrt.negative", unary_op_pattern("negative")), + ("tensorrt.nn.batch_flatten", unary_op_pattern("nn.batch_flatten")), + ("tensorrt.sin", unary_op_pattern("sin"), trt_version_annotate_fn((5, 1, 5))), + ("tensorrt.clip", unary_op_pattern("clip")), + ("tensorrt.cos", unary_op_pattern("cos"), trt_version_annotate_fn((5, 1, 5))), + ("tensorrt.atan", unary_op_pattern("atan"), trt_version_annotate_fn((5, 1, 5))), + ("tensorrt.ceil", unary_op_pattern("ceil"), trt_version_annotate_fn((5, 1, 5))), + ("tensorrt.floor", unary_op_pattern("floor")), + ("tensorrt.erf", unary_op_pattern("erf"), trt_version_annotate_fn((7, 0, 0))), + ("tensorrt.sum", unary_op_pattern("sum"), reduce_annotate_fn), + ("tensorrt.prod", unary_op_pattern("prod"), reduce_annotate_fn), + ("tensorrt.max", unary_op_pattern("max"), reduce_annotate_fn), + ("tensorrt.min", unary_op_pattern("min"), reduce_annotate_fn), + ("tensorrt.max", unary_op_pattern("max"), reduce_annotate_fn), + ("tensorrt.concatenate", unary_op_pattern("concatenate"), concatenate_annotate_fn), + ("tensorrt.expand_dims", unary_op_pattern("expand_dims"), expand_dims_annotate_fn), + ( + "tensorrt.layout_transform", + unary_op_pattern("layout_transform"), + layout_transform_annotate_fn, + ), + ("tensorrt.transpose", unary_op_pattern("transpose"), transpose_annotate_fn), + ("tensorrt.reshape", unary_op_pattern("reshape"), reshape_annotate_fn), + ("tensorrt.split", unary_op_pattern("split"), split), + ("tensorrt.nn.pad", unary_op_pattern("nn.pad"), pad_annotate_fn), + ("tensorrt.strided_slice", unary_op_pattern("strided_slice"), strided_slice_annotate_fn), + ( + "tensorrt.nn.adaptive_avg_pool2d", + unary_op_pattern("nn.adaptive_avg_pool2d"), + adaptive_avg_pool2d_annotate_fn, + ), + ("tensorrt.nn.max_pool3d", unary_op_pattern("nn.max_pool3d"), max_pool_3d_annotate_fn), + ("tensorrt.nn.avg_pool3d", unary_op_pattern("nn.avg_pool3d"), avg_pool_3d_annotate_fn), + ( + "tensorrt.nn.conv3d_transpose", + unary_op_pattern("nn.conv3d_transpose"), + conv3d_transpose_annotate_fn, + ), + ("tensorrt.nn.softmax", unary_op_pattern("nn.softmax"), softmax_annotate_fn), + ("tensorrt.nn.layer_norm", unary_op_pattern("nn.layer_norm"), layer_norm_annotate_fn), + ("tensorrt.nn.max_pool2d", unary_op_pattern("nn.max_pool2d"), max_pool_2d_annotate_fn), + ("tensorrt.nn.avg_pool2d", unary_op_pattern("nn.avg_pool2d"), avg_pool_2d_annotate_fn), + ("tensorrt.nn.max_pool3d", unary_op_pattern("nn.max_pool3d"), max_pool_3d_annotate_fn), + ( + "tensorrt.nn.global_max_pool2d", + unary_op_pattern("nn.global_max_pool2d"), + global_max_pool_2d_annotate_fn, + ), + ( + "tensorrt.nn.global_avg_pool2d", + unary_op_pattern("nn.global_avg_pool2d"), + global_avg_pool_2d_annotate_fn, + ), + ] + + class IsComputeIntensiveGraph(ExprVisitor): """ Visits the Graph recursively and checks if it contains compute heavy ops like convolutions and @@ -1175,8 +1146,7 @@ def visit_call(self, call): subgraphs_to_remove.append(name) # Create new pruned module new_mod = tvm.IRModule(mod.functions, mod.type_definitions) - new_mod["main"] = SubgraphRemover( - subgraphs_to_remove, mod, new_mod).visit(mod["main"]) + new_mod["main"] = SubgraphRemover(subgraphs_to_remove, mod, new_mod).visit(mod["main"]) new_mod = transform.RemoveUnusedFunctions()(new_mod) return new_mod diff --git a/python/tvm/relay/transform/transform.py b/python/tvm/relay/transform/transform.py index ac0f61a189b5..e4ee14b62941 100644 --- a/python/tvm/relay/transform/transform.py +++ b/python/tvm/relay/transform/transform.py @@ -543,10 +543,7 @@ def MergeComposite(pattern_table): for tup in pattern_table: if len(tup) == 2: pattern_name, pattern = tup - - def check(extract): - return True - + check = lambda extract: True elif len(tup) == 3: pattern_name, pattern, check = tup @@ -792,17 +789,22 @@ def Inline(): return _ffi_api.Inline() -def UnmergeComposites(): - """Perform inlining on the given Relay IR module. The global functions that - are marked as `inline` should be always inlined. A cost model will be - needed in the future to decide if it is profitable to inline the function. +def InlineComposites(target): + """Perform inlining on the given Relay IR module. The functions originate + from the MergeComposite pass based on an input pattern table will fold back + to main. Currently, this is used for the TRT BYOC which expects a single + primitive function to operate on. + Parameters + ---------- + target: str + The byoc target for which ops need to fold back to primitive function. Returns ------- ret: tvm.transform.Pass The registered pass that performs inlining for a Relay IR module. """ - return _ffi_api.UnmergeComposites() + return _ffi_api.InlineComposites(target) def gradient(expr, mod=None, mode="higher_order"): diff --git a/src/relay/transforms/inline_composites.cc b/src/relay/transforms/inline_composites.cc new file mode 100644 index 000000000000..b878bb247874 --- /dev/null +++ b/src/relay/transforms/inline_composites.cc @@ -0,0 +1,119 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +/*! + * \file src/relay/transforms/inline_composites.cc + * \brief Undo the partioned graphs originate from merge composite. + */ +#include +#include +#include + +#include "../analysis/call_graph.h" +#include "../op/call/call.h" + +using namespace tvm::runtime; + +namespace tvm { + +namespace relay { + +class Unmerger : ExprMutator { + public: + explicit Unmerger(CallGraphEntry* cur_node, CallGraphNode* call_graph) + : cur_node_(cur_node), call_graph_(call_graph) {} + + Expr VisitExpr_(const CallNode* call_node) final { + Call vanilla_call = GetAnyCall(call_node); + const auto* function_node = vanilla_call->op.as(); + + if (function_node) { + Array new_args; + new_args.reserve(vanilla_call->args.size()); + for (auto arg : vanilla_call->args) { + new_args.push_back(VisitExpr(arg)); + } + + Map bind_map; + for (size_t i = 0; i < new_args.size(); i++) { + bind_map.Set(function_node->params[i], new_args[i]); + } + + // Attrs need to be empty at this point to avoid propagating Composite and + // PartitionedFromPattern that fiddling TRT code gen for registered ops. + return Bind(function_node->body, bind_map); + } + + return ExprMutator::VisitExpr_(call_node); + } + + Function Unmerge(const Function& func) { + return WithFields(func, func->params, VisitExpr(func->body)); + } + + private: + /*! + * \brief The current call graph entry that is being handled. Each entry + * contains a global function. + */ + CallGraphEntry* cur_node_; + /*! \brief The call graph that is used for global function lookup. */ + const CallGraphNode* call_graph_; +}; + +IRModule InlineComposites(const IRModule& module, runtime::String target) { + CallGraph cg(module); + auto topo = cg->TopologicalOrder(); + std::reverse(topo.begin(), topo.end()); + std::unordered_set original_entry; + ICHECK(target.defined()); + for (auto* it : topo) { + auto base_func = module->Lookup(it->GetNameHint()); + + if (!base_func->GetAttr(attr::kCompiler).defined() && + base_func->GetAttr(attr::kCompiler) != target) { + return module; + } + + if (it->GetNameHint() != "main") { + if (const auto* fn = base_func.as()) { + auto func = GetRef(fn); + auto new_func = Unmerger(it, cg.operator->()).Unmerge(func); + cg->module->Update(it->GetGlobalVar(), new_func); + } + } + } + return module; +} + +namespace transform { + +Pass InlineComposites(runtime::String target) { + runtime::TypedPackedFunc pass_func = + [=](IRModule m, PassContext pc) { return relay::InlineComposites(m, target); }; + return CreateModulePass(pass_func, 0, "InlineComposites", {}); +} + +TVM_REGISTER_GLOBAL("relay._transform.InlineComposites").set_body_typed(InlineComposites); + +} // namespace transform + +} // namespace relay + +} // namespace tvm diff --git a/src/relay/transforms/unmerge_composites.cc b/src/relay/transforms/unmerge_composites.cc deleted file mode 100644 index 72c3ec6978ab..000000000000 --- a/src/relay/transforms/unmerge_composites.cc +++ /dev/null @@ -1,254 +0,0 @@ -/* - * Licensed to the Apache Software Foundation (ASF) under one - * or more contributor license agreements. See the NOTICE file - * distributed with this work for additional information - * regarding copyright ownership. The ASF licenses this file - * to you under the Apache License, Version 2.0 (the - * "License"); you may not use this file except in compliance - * with the License. You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, - * software distributed under the License is distributed on an - * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY - * KIND, either express or implied. See the License for the - * specific language governing permissions and limitations - * under the License. - */ - -/*! - * \file src/relay/transforms/unmerge_composites.cc - * \brief Unmerges composite functions - */ -#include -#include -#include - -#include "../analysis/call_graph.h" -#include "../op/call/call.h" - -using namespace tvm::runtime; - -namespace tvm { - -namespace relay { - -class Unmerger : ExprMutator { - public: - explicit Unmerger(CallGraphEntry* cur_node, CallGraphNode* call_graph) - : cur_node_(cur_node), call_graph_(call_graph) {} - - Expr VisitExpr_(const CallNode* call_node) final { - // We can work with calls in both pre- and post-lowered form. - Call vanilla_call = GetAnyCall(call_node); - VLOG(1) << "Vanilla call " << vanilla_call->op << std::endl; - VLOG(1) << "Vanilla call " << vanilla_call->op->checked_type_ << std::endl; - // VLOG(1) << "Vanilla attrs " << vanilla_call->attrs << std::endl; - const auto* global_var_node = vanilla_call->op.as(); - const auto* function_var_node = vanilla_call->op.as(); - // const auto* function__node = vanilla_call->op.as(); - - if (global_var_node) { - VLOG(1) << "Global Var node"; - } - - if (function_var_node) { - VLOG(1) << "For existing function "; - Function gv = GetRef(function_var_node); - const auto* fn = gv.as(); - ICHECK(fn) << "Expected to work on a Relay function."; - // auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, fn->attrs); - - Array new_args; - new_args.reserve(vanilla_call->args.size()); - for (auto arg : vanilla_call->args) { - new_args.push_back(VisitExpr(arg)); - } - - Map bind_map; - for (size_t i = 0; i < new_args.size(); i++) { - bind_map.Set(fn->params[i], new_args[i]); - } - - auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, {}); - VLOG(1) << "Params :" << func->params; - VLOG(1) << "Ret type :" << func->ret_type; - VLOG(1) << "Type Params :" << func->type_params; - VLOG(1) << "Attrs :" << func->attrs; - VLOG(1) << "Func body:" << func->body; - - return Bind(func->body, bind_map); - - // return func; - // return func->body; - // VLOG(1) << "gv " << func; - // return func->body; - // auto base_func = call_graph_->GetGlobalFunction(global); - } - // ICHECK(function_var_node); - // VLOG(1) << std::endl << "function : " << function__node << std::endl; - // VLOG(1) << "Function var node : " << function_var_node->attrs; - // if(function_var_node) { - // VLOG(1) << "As function var node"; - // } - // if (global_var_node) { - // ICHECK(function_var_node); - // ICHECK(global_var_node); - // GlobalVar gv = GetRef(global_var_node); - // Function gv = GetRef(function_var_node); - - // auto* cg_node = (*call_graph_)[gv->name_hint]; - // // if (CanInline(cg_node)) { - // Array new_args; - // new_args.reserve(vanilla_call->args.size()); - // for (auto arg : vanilla_call->args) { - // new_args.push_back(VisitExpr(arg)); - // } - // cur_node_->RemoveCallTo(gv); - // return MakeNewExpr(gv, new_args, GetRef(call_node)); - // } - // else: fallthrough - // } - // else: fallthrough - - // If not calling a global function then nothing to inline. - return ExprMutator::VisitExpr_(call_node); - } - - Expr VisitExpr_(const GlobalVarNode* gvn) final { - GlobalVar gv = GetRef(gvn); - auto* cg_node = (*call_graph_)[gv->name_hint]; - if (CanInline(cg_node)) { - cur_node_->RemoveCallTo(gv); - return MakeNewExpr(gv, {}, GetRef(gvn)); - } - return ExprMutator::VisitExpr_(gvn); - } - - Function Unmerge(const Function& func) { - return WithFields(func, func->params, VisitExpr(func->body)); - } - - private: - bool CanInline(const CallGraphEntry* cg_node) { - // The node must be a leaf node and it cannot be recursive. - return true; - if (!cg_node->empty() || cg_node->IsRecursive()) return false; - - auto base_func = call_graph_->GetGlobalFunction(cg_node->GetGlobalVar()); - const auto* function_node = base_func.as(); - if (!function_node) { - // Can't inline PrimFuncs! - return false; - } - // The body of a global functions must be defined. - if (!function_node->body.defined()) return false; - - // The function must be annotated with the inline attribute. - // (Note that external functions do not have this attribute!) - // if (!function_node->HasNonzeroAttr(attr::kInline)) return false; - if (!function_node->HasNonzeroAttr(attr::kInline)) return false; - - // The function is not able to be inlined if any callee under the CallGraph - // of this function cannot be inlined. - for (const auto& it : *cg_node) { - if (!CanInline(it.second)) { - return false; - } - } - - return true; - } - - // Make a new Relay expression to replace \p expr. - Expr MakeNewExpr(const GlobalVar& global, const Array& args, const Expr& expr) { - ICHECK(expr->IsInstance() || expr->IsInstance()); - auto base_func = call_graph_->GetGlobalFunction(global); - const auto* fn = base_func.as(); - VLOG(1) << "Make new expr " << fn; - ICHECK(fn) << "Expected to work on a Relay function."; - - // There is an inconsistency here, the function itself gets shallow-copied but the body is not - // shallow-copied. - auto func = Function(fn->params, fn->body, fn->ret_type, fn->type_params, fn->attrs); - // Inline the function body to the caller if this function uses default - // compiler, i.e. no external codegen is needed. - if (!func->GetAttr(attr::kCompiler).defined() && - !func->GetAttr(attr::kExternalSymbol).defined()) { - ICHECK_EQ(func->params.size(), args.size()) - << "Mismatch found in the number of parameters and call args"; - // Bind the parameters with call args. - Map bind_map; - for (size_t i = 0; i < args.size(); i++) { - bind_map.Set(fn->params[i], args[i]); - } - if (const auto* gvn = expr.as()) { - auto ret_type = gvn->checked_type(); - // Cannot replace TensorType/TensorTupleType with FuncType. Therefore, - // we simply inline the function as a closure instead of directly using - // its body when the global var returns FuncType. - return ret_type->IsInstance() ? std::move(func) : func->body; - } else { - ICHECK(expr->IsInstance()); - return Bind(func->body, bind_map); - } - } else if (const auto* call_node = expr.as()) { - return Call(func, args, call_node->attrs, call_node->type_args); - } else { - return std::move(func); - } - } - - /*! - * \brief The current call graph entry that is being handled. Each entry - * contains a global function. - */ - CallGraphEntry* cur_node_; - /*! \brief The call graph that is used for global function lookup. */ - const CallGraphNode* call_graph_; -}; - -IRModule UnmergeComposites(const IRModule& module) { - CallGraph cg(module); - auto topo = cg->TopologicalOrder(); - std::reverse(topo.begin(), topo.end()); - std::unordered_set original_entry; - VLOG_CONTEXT << "Unmerge Composite"; - - VLOG(1) << "Topo size: " << topo.size(); - for (auto* it : topo) { - auto base_func = module->Lookup(it->GetNameHint()); - if (it->GetNameHint() != "main") { - // Check kSymbol that is the correct one - // Check kCompiler that is the correct one - if (const auto* fn = base_func.as()) { - VLOG(1) << "Func name " << it->GetNameHint() << std::endl << "-------" << std::endl; - auto func = GetRef(fn); - auto new_func = Unmerger(it, cg.operator->()).Unmerge(func); - - cg->module->Update(it->GetGlobalVar(), new_func); - } - } - } - VLOG(1) << "Post unmerge module " << std::endl; - VLOG(1) << module; - VLOG(1) << "------------- " << std::endl; - return module; -} - -namespace transform { - -Pass UnmergeComposites() { - runtime::TypedPackedFunc pass_func = - [=](IRModule m, PassContext pc) { return relay::UnmergeComposites(m); }; - return CreateModulePass(pass_func, 1, "UnmergeComposites", {}); -} - -TVM_REGISTER_GLOBAL("relay._transform.UnmergeComposites").set_body_typed(UnmergeComposites); - -} // namespace transform - -} // namespace relay - -} // namespace tvm \ No newline at end of file From a3dc5545f80b02cb8392f30bb4c60cf7b784b900 Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Fri, 25 Mar 2022 02:20:30 -0700 Subject: [PATCH 6/8] Address PR comments --- python/tvm/relay/op/contrib/tensorrt.py | 90 +++++++++++++++++------ src/relay/transforms/inline_composites.cc | 14 ++-- 2 files changed, 74 insertions(+), 30 deletions(-) diff --git a/python/tvm/relay/op/contrib/tensorrt.py b/python/tvm/relay/op/contrib/tensorrt.py index eb814ab5012a..f24d366e598a 100644 --- a/python/tvm/relay/op/contrib/tensorrt.py +++ b/python/tvm/relay/op/contrib/tensorrt.py @@ -106,6 +106,7 @@ def partition_for_tensorrt( max_workspace_size=1 << 30, use_fp16=False, use_uint8=False, + use_patterns=False, ): """Partition the graph greedily offloading supported operators to TensorRT. @@ -136,6 +137,9 @@ def partition_for_tensorrt( lower runtime, or if no low-precision implementation exists. use_uint8: Optional[bool] Allows, TRT to automatically convert FP32 inputs to UINT8. + use_patterns: Optional[bool] + Switches to use pattern-based op suppot by applying MergeCompsite and InlineComposites + passes. Returns ------- mod_and_config : Tuple[Module, Dict[str, Any]] @@ -164,34 +168,74 @@ def partition_for_tensorrt( if params: mod["main"] = bind_params_by_name(mod["main"], params) - seq = tvm.transform.Sequential( - [ - transform.InferType(), - RemoveDropoutPass(), - transform.RemoveUnusedFunctions(), - transform.ConvertLayout( - { - "nn.conv1d": ["NCW", "default"], - "nn.conv2d": ["NCHW", "default"], - "nn.conv3d": ["NCDHW", "default"], - "nn.conv2d_transpose": ["NCHW", "default"], - } - ), - transform.FoldConstant(), - transform.MergeComposite(pattern_table()), - transform.AnnotateTarget("tensorrt"), - transform.MergeCompilerRegions(), - transform.PartitionGraph(), - transform.InlineComposites("tensorrt"), - transform.InferType(), - ] - ) + + seq = get_pass_order(use_patterns) with tvm.transform.PassContext(opt_level=3, config={"relay.ext.tensorrt.options": config}): mod = seq(mod) mod = prune_tensorrt_subgraphs(mod) return mod, config +def get_pass_order(use_patterns): + """ + Get the pass ordering based on using predicates or patterns. + + Parameters + ---------- + use_patterns: Bool + True if pass needs to work with op patterns + Returns + ---------- + ret : Sequential + Pass object + """ + return ( + tvm.transform.Sequential( + [ + transform.InferType(), + RemoveDropoutPass(), + transform.RemoveUnusedFunctions(), + transform.ConvertLayout( + { + "nn.conv1d": ["NCW", "default"], + "nn.conv2d": ["NCHW", "default"], + "nn.conv3d": ["NCDHW", "default"], + "nn.conv2d_transpose": ["NCHW", "default"], + } + ), + transform.FoldConstant(), + transform.MergeComposite(pattern_table()), + transform.AnnotateTarget("tensorrt"), + transform.MergeCompilerRegions(), + transform.PartitionGraph(), + transform.InlineComposites("tensorrt"), + transform.InferType(), + ] + ) + if use_patterns + else tvm.transform.Sequential( + [ + transform.InferType(), + RemoveDropoutPass(), + transform.RemoveUnusedFunctions(), + transform.ConvertLayout( + { + "nn.conv1d": ["NCW", "default"], + "nn.conv2d": ["NCHW", "default"], + "nn.conv3d": ["NCDHW", "default"], + "nn.conv2d_transpose": ["NCHW", "default"], + } + ), + transform.FoldConstant(), + transform.AnnotateTarget("tensorrt"), + transform.MergeCompilerRegions(), + transform.PartitionGraph(), + transform.InferType(), + ] + ) + ) + + def check_dynamism(args, op_name): """ Check for dynamism inside any of the args in the op. @@ -451,7 +495,7 @@ def dense_annotate_fn(expr): # pylint: disable=unused-variable @_register_external_dynamic_check_func("nn.batch_matmul") -def batch_matmul_annotate_fn(expr): # pylint: disable=unused-variable +def batch_matmul_annotate_fn(expr): """Check if dense is supported by TensorRT.""" args = expr.args diff --git a/src/relay/transforms/inline_composites.cc b/src/relay/transforms/inline_composites.cc index b878bb247874..63e7d078b0c5 100644 --- a/src/relay/transforms/inline_composites.cc +++ b/src/relay/transforms/inline_composites.cc @@ -34,12 +34,12 @@ namespace tvm { namespace relay { -class Unmerger : ExprMutator { +class CompositeInliner : public MixedModeMutator { public: - explicit Unmerger(CallGraphEntry* cur_node, CallGraphNode* call_graph) + explicit CompositeInliner(CallGraphEntry* cur_node, CallGraphNode* call_graph) : cur_node_(cur_node), call_graph_(call_graph) {} - Expr VisitExpr_(const CallNode* call_node) final { + Expr Rewrite_(const CallNode* call_node) { Call vanilla_call = GetAnyCall(call_node); const auto* function_node = vanilla_call->op.as(); @@ -60,10 +60,10 @@ class Unmerger : ExprMutator { return Bind(function_node->body, bind_map); } - return ExprMutator::VisitExpr_(call_node); + return MixedModeMutator::VisitExpr_(call_node); } - Function Unmerge(const Function& func) { + Function Inline(const Function& func) { return WithFields(func, func->params, VisitExpr(func->body)); } @@ -88,13 +88,13 @@ IRModule InlineComposites(const IRModule& module, runtime::String target) { if (!base_func->GetAttr(attr::kCompiler).defined() && base_func->GetAttr(attr::kCompiler) != target) { - return module; + continue; } if (it->GetNameHint() != "main") { if (const auto* fn = base_func.as()) { auto func = GetRef(fn); - auto new_func = Unmerger(it, cg.operator->()).Unmerge(func); + auto new_func = CompositeInliner(it, cg.operator->()).Inline(func); cg->module->Update(it->GetGlobalVar(), new_func); } } From 2a6060f779f4034e8f8f801f3524581a6722359e Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Thu, 31 Mar 2022 03:06:48 -0700 Subject: [PATCH 7/8] Add unit tests for inline composites pass --- .../relay/test_pass_inline_composites.py | 217 ++++++++++++++++++ 1 file changed, 217 insertions(+) create mode 100644 tests/python/relay/test_pass_inline_composites.py diff --git a/tests/python/relay/test_pass_inline_composites.py b/tests/python/relay/test_pass_inline_composites.py new file mode 100644 index 000000000000..d2fb9fa1a6e3 --- /dev/null +++ b/tests/python/relay/test_pass_inline_composites.py @@ -0,0 +1,217 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +# pylint: disable=invalid-name, missing-docstring, too-many-statements +"""Unit tests for inline composites.""" +import pytest +import tvm +from tvm import relay, tir +from tvm.relay.dataflow_pattern import TupleGetItemPattern, is_op, wildcard +from tvm.relay.testing import run_opt_pass + +""" +The inline composite pass is designed to inline multiple kernel generated through +the merge composite composite pass. The underlying idea is to inline N kernels +produced from merge composite based on a given set of pattern into a single IR module. +Also, clears Composite and PartionedFromPatterns that infer with certain BYOC implementations + +For example suppose we have the graph: + + a b + \ / + add + | + relu + +Merge composite will wrap each standalone op to it's own function, while setting Composite and +PartitionedFromPattern attrs. + +Relay IR after merge composite pass when registering each op as a standalone pattern: +fn (%a: Tensor[(10, 10), float32], %b: Tensor[(10, 10), float32]) -> Tensor[(10, 10), float32] { + %0 = fn (%FunctionVar_0_01: Tensor[(10, 10), float32], %FunctionVar_0_1: Tensor[(10, 10), float32], PartitionedFromPattern="add_", Composite="add") -> Tensor[(10, 10), float32] { + add(%FunctionVar_0_01, %FunctionVar_0_1) /* ty=Tensor[(10, 10), float32] */ + }; + %1 = %0(%a, %b) /* ty=Tensor[(10, 10), float32] */; + %2 = fn (%FunctionVar_0_0: Tensor[(10, 10), float32], PartitionedFromPattern="nn.relu_", Composite="nn.relu") -> Tensor[(10, 10), float32] { + nn.relu(%FunctionVar_0_0) /* ty=Tensor[(10, 10), float32] */ + }; + %2(%1) /* ty=Tensor[(10, 10), float32] */ +} + +Relay IR after inline composites pass: +fn (%a: Tensor[(10, 10), float32], %b: Tensor[(10, 10), float32]) -> Tensor[(10, 10), float32] { + %0 = add(%a, %b) /* ty=Tensor[(10, 10), float32] */; + nn.relu(%0) /* ty=Tensor[(10, 10), float32] */ +} + +One conventient use of this pass is to use Pattern-based operator support to move away +from the original operator predicates, and inline them into a single primitive function to offload it +to an external BYOC backend, such as TensorRT. +""" + + +def make_conv_bias_relu_pattern(): + r"""Create a pattern to match the following graph. + + conv2d + | + bias_add + | + relu + """ + x = wildcard() + y = wildcard() + z = wildcard() + conv_node = is_op("nn.conv2d")(x, y) + bias_node = is_op("nn.bias_add")(conv_node, z) + r = is_op("nn.relu")(bias_node) + return r + + +def make_add_relu_pattern(): + r"""Create a pattern to match the following graph. + + add + | + relu + """ + add_node = wildcard() + wildcard() + r = is_op("nn.relu")(add_node) + return r + + +def make_relu_pattern(): + r"""Create a pattern to match the following graph + a + | + relu + | + """ + pattern = is_op("nn.relu")(wildcard()) + return pattern + + +def make_add_pattern(): + r"""Create a pattern to match the following graph + a b + \ / + add + | + """ + pattern = is_op("add")(wildcard(), wildcard()) + return pattern + + +def check_result(pattern_table, graph, expected_graph, import_prelude=False): + """Utility function to check inline composites results.""" + result = run_opt_pass( + graph, relay.transform.MergeComposite(pattern_table), import_prelude=import_prelude + ) + print("merge composite reusult") + print(result) + print("---------------------") + result = run_opt_pass( + graph, relay.transform.InlineComposites(target=""), import_prelude=import_prelude + ) + print(result) + print("-----------------") + assert not relay.analysis.free_vars(result), "Found free vars in the result graph: {0}".format( + str(result) + ) + expected = run_opt_pass(expected_graph, relay.transform.InferType()) + print(expected) + print("---------------") + assert tvm.ir.structural_equal( + result, expected, map_free_vars=True + ), "Graph mismatch: output vs. expected\n{0}\n=====\n{1}".format(str(result), str(expected)) + + +def test_single_op_registry(): + r"""Test inline composite pass is correctly inline the merge composite result. + + We could expect the pattern `make_add_relu_pattern` to be merged + into a single op `add_relu`. + + a b a b + \ / a b \ / + add ====> \ / ===> add + | add_relu | + relu relu + + """ + pattern_table = [("add", make_add_pattern()), ("nn.relu", make_relu_pattern())] + + def before(): + a = relay.var("a", shape=(10, 10)) + b = relay.var("b", shape=(10, 10)) + add_node = relay.add(a, b) + r = relay.nn.relu(add_node) + return relay.Function([a, b], r) + + def expected(): + a = relay.var("a", shape=(10, 10)) + b = relay.var("b", shape=(10, 10)) + + # add_relu function + in_1 = relay.var("in_1", shape=(10, 10)) + in_2 = relay.var("in_2", shape=(10, 10)) + add_node = relay.add(in_1, in_2) + relu_node = relay.nn.relu(add_node) + add_relu = relay.Function([in_1, in_2], relu_node) + return add_relu + + check_result(pattern_table, before(), expected()) + + +def test_simple_merge(): + r"""Test inline composite pass is correctly inline the merge composite result. + + We could expect the pattern `make_add_relu_pattern` to be merged + into a single op `add_relu`. + + a b a b + \ / a b \ / + add ====> \ / ===> add + | add_relu | + relu relu + + """ + pattern_table = [("add_relu", make_add_relu_pattern())] + + def before(): + a = relay.var("a", shape=(10, 10)) + b = relay.var("b", shape=(10, 10)) + add_node = relay.add(a, b) + r = relay.nn.relu(add_node) + return relay.Function([a, b], r) + + def expected(): + a = relay.var("a", shape=(10, 10)) + b = relay.var("b", shape=(10, 10)) + + # add_relu function + in_1 = relay.var("in_1", shape=(10, 10)) + in_2 = relay.var("in_2", shape=(10, 10)) + add_node = relay.add(in_1, in_2) + relu_node = relay.nn.relu(add_node) + add_relu = relay.Function([in_1, in_2], relu_node) + return add_relu + + check_result(pattern_table, before(), expected()) + + +if __name__ == "__main__": + pytest.main() From 6c39b810190b1b91fb846481aad82f7b4b0fa5d3 Mon Sep 17 00:00:00 2001 From: Michalis Papapdimitriou Date: Fri, 1 Apr 2022 05:01:10 -0700 Subject: [PATCH 8/8] Address PR comments --- .../relay/test_pass_inline_composites.py | 94 +++++-------------- 1 file changed, 21 insertions(+), 73 deletions(-) diff --git a/tests/python/relay/test_pass_inline_composites.py b/tests/python/relay/test_pass_inline_composites.py index d2fb9fa1a6e3..54fc08c87918 100644 --- a/tests/python/relay/test_pass_inline_composites.py +++ b/tests/python/relay/test_pass_inline_composites.py @@ -57,30 +57,12 @@ nn.relu(%0) /* ty=Tensor[(10, 10), float32] */ } -One conventient use of this pass is to use Pattern-based operator support to move away +One convenient use of this pass is to use Pattern-based operator support to move away from the original operator predicates, and inline them into a single primitive function to offload it to an external BYOC backend, such as TensorRT. """ -def make_conv_bias_relu_pattern(): - r"""Create a pattern to match the following graph. - - conv2d - | - bias_add - | - relu - """ - x = wildcard() - y = wildcard() - z = wildcard() - conv_node = is_op("nn.conv2d")(x, y) - bias_node = is_op("nn.bias_add")(conv_node, z) - r = is_op("nn.relu")(bias_node) - return r - - def make_add_relu_pattern(): r"""Create a pattern to match the following graph. @@ -115,57 +97,40 @@ def make_add_pattern(): return pattern -def check_result(pattern_table, graph, expected_graph, import_prelude=False): +def check_success_composite_pass(func): + return func.body.op.attrs["Composite"] is not None + + +def check_result(pattern_table, expected_graph, import_prelude=False): """Utility function to check inline composites results.""" result = run_opt_pass( - graph, relay.transform.MergeComposite(pattern_table), import_prelude=import_prelude + expected_graph, relay.transform.MergeComposite(pattern_table), import_prelude=import_prelude ) - print("merge composite reusult") - print(result) - print("---------------------") + assert check_success_composite_pass( + result + ), "Merge Composite pass didn't produced partioned from Pattern" result = run_opt_pass( - graph, relay.transform.InlineComposites(target=""), import_prelude=import_prelude + expected_graph, relay.transform.InlineComposites(target=""), import_prelude=import_prelude ) - print(result) - print("-----------------") assert not relay.analysis.free_vars(result), "Found free vars in the result graph: {0}".format( str(result) ) expected = run_opt_pass(expected_graph, relay.transform.InferType()) - print(expected) - print("---------------") assert tvm.ir.structural_equal( result, expected, map_free_vars=True ), "Graph mismatch: output vs. expected\n{0}\n=====\n{1}".format(str(result), str(expected)) def test_single_op_registry(): - r"""Test inline composite pass is correctly inline the merge composite result. - - We could expect the pattern `make_add_relu_pattern` to be merged - into a single op `add_relu`. + r"""Test inline composite pass is correctly inline the post-merge composite graph. - a b a b - \ / a b \ / - add ====> \ / ===> add - | add_relu | - relu relu + We could expect the patterns `make_add_pattern` and `make_relu_pattern` to be inlined + into a single func instead of an single func per registered pattern. """ pattern_table = [("add", make_add_pattern()), ("nn.relu", make_relu_pattern())] - def before(): - a = relay.var("a", shape=(10, 10)) - b = relay.var("b", shape=(10, 10)) - add_node = relay.add(a, b) - r = relay.nn.relu(add_node) - return relay.Function([a, b], r) - def expected(): - a = relay.var("a", shape=(10, 10)) - b = relay.var("b", shape=(10, 10)) - - # add_relu function in_1 = relay.var("in_1", shape=(10, 10)) in_2 = relay.var("in_2", shape=(10, 10)) add_node = relay.add(in_1, in_2) @@ -173,30 +138,12 @@ def expected(): add_relu = relay.Function([in_1, in_2], relu_node) return add_relu - check_result(pattern_table, before(), expected()) - - -def test_simple_merge(): - r"""Test inline composite pass is correctly inline the merge composite result. + check_result(pattern_table, expected()) - We could expect the pattern `make_add_relu_pattern` to be merged - into a single op `add_relu`. - a b a b - \ / a b \ / - add ====> \ / ===> add - | add_relu | - relu relu - - """ - pattern_table = [("add_relu", make_add_relu_pattern())] - - def before(): - a = relay.var("a", shape=(10, 10)) - b = relay.var("b", shape=(10, 10)) - add_node = relay.add(a, b) - r = relay.nn.relu(add_node) - return relay.Function([a, b], r) +def test_mix_fused_and_single_op(): + r"""Test inline composite pass is correctly inline the merge composite result""" + pattern_table = [("add_relu", make_add_relu_pattern()), ("nn.relu", make_relu_pattern())] def expected(): a = relay.var("a", shape=(10, 10)) @@ -207,10 +154,11 @@ def expected(): in_2 = relay.var("in_2", shape=(10, 10)) add_node = relay.add(in_1, in_2) relu_node = relay.nn.relu(add_node) - add_relu = relay.Function([in_1, in_2], relu_node) + relu_nd = relay.nn.relu(relu_node) + add_relu = relay.Function([in_1, in_2], relu_nd) return add_relu - check_result(pattern_table, before(), expected()) + check_result(pattern_table, expected()) if __name__ == "__main__":