From 5ffc51c5832bc8f692441c0b5bbf044f8d29b2ca Mon Sep 17 00:00:00 2001 From: Masahiro Masuda Date: Wed, 15 Dec 2021 13:29:42 +0900 Subject: [PATCH 1/2] Add cutlass conv2d activation (bias, relu, sigmoid) commit e4e273ae74a8e54ab1ae1414ce9b6bfcc2b3d530 Merge: 0489d1418 77c938550 Author: Masahiro Masuda Date: Mon Dec 13 11:58:54 2021 +0900 Merge branch 'partition-constant-unbind' into cutlass-conv2d-fusion commit 77c9385501595f804bd33b436aed0cc192059a10 Author: Masahiro Masuda Date: Mon Dec 13 11:58:18 2021 +0900 add test commit ab01b3aae36ef88ec299ff5e9d1bbdf5fd7268f6 Author: Masahiro Masuda Date: Mon Dec 13 11:55:06 2021 +0900 make constant binding in PartitionGraph optional commit 0489d1418c5f49b95ad220b8f0e57ba48443c6a3 Author: Masahiro Masuda Date: Sun Dec 12 21:52:29 2021 +0900 support sigmoid fusion (only fp32 accum for now) commit 3705bbd6b77ec8d8c416ab00738fe8b9c91e1535 Author: Masahiro Masuda Date: Sun Dec 12 20:50:58 2021 +0900 conv2d fusion test worked commit 05b51c9f94c9a4f31afeb126c8749675be423892 Author: Masahiro Masuda Date: Sun Dec 12 20:34:10 2021 +0900 fix bias stride commit 7cf40e719a8c1c464e2ee925018d8ed90b145dcd Author: Masahiro Masuda Date: Sun Dec 12 20:01:21 2021 +0900 use nobetascaling commit 274ec028845e296b012c23830a0be6f48cbee767 Author: Masahiro Masuda Date: Sun Dec 12 19:12:58 2021 +0900 adding fusion support to codegen commit 0de5ebdb2e318129a1a4d3a2e568849a993525c4 Author: Masahiro Masuda Date: Sun Dec 12 18:39:08 2021 +0900 partition working commit c08bb38e3eea90168dc036cb4cbfbcdef288bac6 Author: Masahiro Masuda Date: Sun Dec 12 17:24:42 2021 +0900 update test commit 81bf9e600b1bbdbc27e0e8b54496fc09b33e0624 Author: Masahiro Masuda Date: Fri Dec 10 13:23:39 2021 +0900 add fused conv2d pattern commit 1c0bbb297da43dab75ad995afbcacd59e9fe4c87 Author: Masahiro Masuda Date: Sun Dec 12 18:29:03 2021 +0900 fix lint commit 463574ce087ca0444b23d4f47baf8066b8fbd3df Author: Masahiro Masuda Date: Sun Dec 12 17:28:38 2021 +0900 fixed conv2d check commit 588c5abe15abbf0339a7972d81b8235bc7460620 Author: Masahiro Masuda Date: Sun Dec 12 15:05:27 2021 +0900 update test commit a447b57ade7c99da4efd38e1bdfc213a64a80fd2 Author: Masahiro Masuda Date: Sun Dec 12 14:54:52 2021 +0900 speed up profiling by removing initialization commit 93cd039ba04dd80e887ec1a71f358cd86a1e5221 Author: Masahiro Masuda Date: Sun Dec 12 08:26:29 2021 +0900 fixed nhwc cudnn depthwise conv commit 6db71727f553ee2009c9d3feb2b019c24459f4d2 Author: Masahiro Masuda Date: Sat Dec 11 15:39:05 2021 +0900 add cache commit f7d17a116acd80c57dfa04aa86577fa30898908d Author: Masahiro Masuda Date: Sat Dec 11 15:05:38 2021 +0900 removed im2col profiling for conv2d commit b724f446d07030e45f30474a1e3c124f1541b496 Author: Masahiro Masuda Date: Fri Dec 10 22:57:54 2021 +0900 black commit fe4687b9f6d41eff66742cdde694680538a3d5e4 Author: Masahiro Masuda Date: Fri Dec 10 22:49:13 2021 +0900 fixed cmd arguement commit ab114f5c3e1a3086255caccd6c0f5df09d7c3755 Author: Masahiro Masuda Date: Fri Dec 10 22:22:19 2021 +0900 conv2d profiler working commit 49ee61f5583f73f0d9b2c18c1861c90fa028f64a Author: Masahiro Masuda Date: Fri Dec 10 20:26:15 2021 +0900 add conv2d profiler commit 49e2c8918a1a2632f60c700675f12f9d83adef24 Author: Masahiro Masuda Date: Sun Dec 12 08:03:36 2021 +0900 do not offload depthwise conv2d commit cd8367768229720020d81d597d751bfa95ec014e Author: Masahiro Masuda Date: Fri Dec 10 13:20:01 2021 +0900 lint fix commit 870823c6d54114896aa9db91eb72c132cdef762f Author: Masahiro Masuda Date: Fri Dec 10 12:54:38 2021 +0900 add comment on IC == 3 case commit 6b780db7f8059a9a58bfc257fbb9cf7cb60b72a2 Author: Masahiro Masuda Date: Fri Dec 10 12:48:33 2021 +0900 check align on N dim commit 308c4dac39f761ac157f032b9fe7028820980294 Author: Masahiro Masuda Date: Fri Dec 10 12:34:42 2021 +0900 fixed check functions for fused cases, run infer type before mergecomposite commit 8d6a1bfee26e9dfdeb1ab001aa89b43b9cd33d74 Author: Masahiro Masuda Date: Fri Dec 10 12:10:59 2021 +0900 test IC=3 convolution commit ffce47de724398222f672b220156b19c8f48a700 Author: Masahiro Masuda Date: Fri Dec 10 12:10:16 2021 +0900 use align1 kernel for unusual channel cases (IC = 3 etc) commit 6cdf205a3451b5fa7ea0cbff0228ae1da775573a Author: Masahiro Masuda Date: Fri Dec 10 12:06:56 2021 +0900 add dtype and layout check in parttern match commit 7743cc6dde442b45d66b4be320268ed24f352422 Author: Masahiro Masuda Date: Fri Dec 10 10:40:53 2021 +0900 add sm75 kernels to sm80 profilings commit efceccb994ec71bc4ba847c791fc4f696eb7f242 Author: Masahiro Masuda Date: Fri Dec 10 10:40:42 2021 +0900 skip legalize when batch size is dynamic commit 65fbc0a0813cb4e980e53240f34f3ba18754ab99 Author: Masahiro Masuda Date: Fri Dec 10 10:36:36 2021 +0900 bug fix in im2col encoding --- python/tvm/contrib/cutlass/build.py | 9 +++ .../tvm/contrib/cutlass/conv2d_operation.py | 35 +++++++--- python/tvm/contrib/cutlass/gen_conv2d.py | 9 +-- python/tvm/contrib/cutlass/gen_gemm.py | 1 - python/tvm/contrib/cutlass/library.py | 2 + python/tvm/relay/op/contrib/cutlass.py | 33 +++++++++- src/relay/backend/contrib/cutlass/codegen.cc | 56 ++++++++++++++-- tests/python/contrib/test_cutlass.py | 66 +++++++++++++++---- 8 files changed, 177 insertions(+), 34 deletions(-) diff --git a/python/tvm/contrib/cutlass/build.py b/python/tvm/contrib/cutlass/build.py index 32a2ebaa2842..a2e6bce8cfea 100644 --- a/python/tvm/contrib/cutlass/build.py +++ b/python/tvm/contrib/cutlass/build.py @@ -87,6 +87,9 @@ def visit_call(self, call): if str(op) == "nn.conv2d": self.op_attrs = call.attrs + for arg in call.args: + self.visit(arg) + def select_gemm_kernel( cutlass_profiler, MM, KK, NN, out_dtype, batched, profile_all, use_multiprocessing @@ -213,6 +216,12 @@ def handle_conv2d( if op_type == "cutlass.conv2d": cutlass_op_def = out["opdef"] + elif op_type == "cutlass.conv2d_bias": + cutlass_op_def = out["opdef_bias"] + elif op_type == "cutlass.conv2d_bias_relu": + cutlass_op_def = out["opdef_bias_relu"] + elif op_type == "cutlass.conv2d_bias_sigmoid": + cutlass_op_def = out["opdef_bias_sigmoid"] else: raise ValueError("%s pattern is not implemented." % op_type) diff --git a/python/tvm/contrib/cutlass/conv2d_operation.py b/python/tvm/contrib/cutlass/conv2d_operation.py index 8a886ff260b8..35308928cdab 100644 --- a/python/tvm/contrib/cutlass/conv2d_operation.py +++ b/python/tvm/contrib/cutlass/conv2d_operation.py @@ -143,6 +143,22 @@ class EmitConv2dInstance: """ Responsible for emitting a CUTLASS template definition.""" def __init__(self): + self.epilogue_default = """ + ${epilogue_functor}< + ${element_c}, + ${epilogue_vector_length}, + ${element_accumulator}, + ${element_epilogue} + >""" + self.epilogue_no_beta_scaling = """ + ${epilogue_functor}< + ${element_c}, + ${epilogue_vector_length}, + ${element_accumulator}, + ${element_epilogue}, + cutlass::epilogue::thread::ScaleType::NoBetaScaling + >""" + self.template = """ // Conv2d${conv_kind_name} ${iterator_algorithm_name} kernel instance "${operation_name}" using ${operation_name} = @@ -159,12 +175,7 @@ def __init__(self): cutlass::gemm::GemmShape<${threadblock_shape_m}, ${threadblock_shape_n}, ${threadblock_shape_k}>, cutlass::gemm::GemmShape<${warp_shape_m}, ${warp_shape_n}, ${warp_shape_k} >, cutlass::gemm::GemmShape<${instruction_shape_m}, ${instruction_shape_n}, ${instruction_shape_k}>, - ${epilogue_functor}< - ${element_c}, - ${epilogue_vector_length}, - ${element_accumulator}, - ${element_epilogue} - >, + ${epilogue}, ${swizzling_functor}, // cutlass::gemm::threadblock::GemmSplitKIdentityThreadblockSwizzle<>, ${stages}, ${math_operator}, @@ -175,7 +186,7 @@ def __init__(self): >::Kernel; """ - def emit(self, operation): + def emit(self, operation, no_beta_scaling=True): """Instantiate a Conv2d kernel from given `operation`.""" warp_shape = [ int( @@ -237,4 +248,12 @@ def emit(self, operation): "align_b": str(operation.B.alignment), } - return substitute_template(self.template, values) + template = substitute_template( + self.template, + { + "epilogue": self.epilogue_no_beta_scaling + if no_beta_scaling + else self.epilogue_default + }, + ) + return substitute_template(template, values) diff --git a/python/tvm/contrib/cutlass/gen_conv2d.py b/python/tvm/contrib/cutlass/gen_conv2d.py index b0d0566d6fab..288f67f39287 100644 --- a/python/tvm/contrib/cutlass/gen_conv2d.py +++ b/python/tvm/contrib/cutlass/gen_conv2d.py @@ -83,15 +83,16 @@ def create_conv2d_operator( op_entry["op"] = op op_entry["src"] = profiler_emitter.emit(op_entry["opdef"], op.procedural_name()) op_entry["name"] = op.procedural_name() - op_entry["runtime"] = 9999999 # fused ops - for epilogue, opdef in zip( + for epilogue, opdef, no_bias_scaling in zip( [ EpilogueFunctor.LinearCombinationBias, EpilogueFunctor.LinearCombinationRelu, + EpilogueFunctor.LinearCombinationSigmoid, ], - ["opdef_bias", "opdef_bias_relu"], + ["opdef_bias", "opdef_bias_relu", "opdef_bias_sigmoid"], + [True, True, False], ): op = Conv2dOperation( ConvKind.Fprop, @@ -107,7 +108,7 @@ def create_conv2d_operator( swizzling_functor_, ) - op_entry[opdef] = kernel_emitter.emit(op) + op_entry[opdef] = kernel_emitter.emit(op, no_bias_scaling) ret.append(op_entry) diff --git a/python/tvm/contrib/cutlass/gen_gemm.py b/python/tvm/contrib/cutlass/gen_gemm.py index c171c5e23a89..7048c32fe1da 100644 --- a/python/tvm/contrib/cutlass/gen_gemm.py +++ b/python/tvm/contrib/cutlass/gen_gemm.py @@ -123,7 +123,6 @@ def create_gemm_operator( DataTypeTag[element_c], op.leading_dim(), ) - op_entry["runtime"] = 9999999 op_entry["tile_description"] = tile_description op_entry["alignment"] = alignment op_entry["data_type"] = data_type diff --git a/python/tvm/contrib/cutlass/library.py b/python/tvm/contrib/cutlass/library.py index 902dc57100a9..8c3f5eb5df63 100644 --- a/python/tvm/contrib/cutlass/library.py +++ b/python/tvm/contrib/cutlass/library.py @@ -148,6 +148,7 @@ class EpilogueFunctor(enum.Enum): LinearCombinationRelu = enum_auto() LinearCombinationBias = enum_auto() LinearCombinationGelu = enum_auto() + LinearCombinationSigmoid = enum_auto() EpilogueFunctorTag = { @@ -155,6 +156,7 @@ class EpilogueFunctor(enum.Enum): EpilogueFunctor.LinearCombinationRelu: "cutlass::epilogue::thread::LinearCombinationRelu", EpilogueFunctor.LinearCombinationBias: "cutlass::epilogue::thread::LinearCombination", EpilogueFunctor.LinearCombinationGelu: "cutlass::epilogue::thread::LinearCombinationGELU", + EpilogueFunctor.LinearCombinationSigmoid: "cutlass::epilogue::thread::LinearCombinationSigmoid", } diff --git a/python/tvm/relay/op/contrib/cutlass.py b/python/tvm/relay/op/contrib/cutlass.py index 0a67581400ed..c706769b2d3d 100644 --- a/python/tvm/relay/op/contrib/cutlass.py +++ b/python/tvm/relay/op/contrib/cutlass.py @@ -57,8 +57,25 @@ def make_batch_matmul_pattern(): return is_op("nn.batch_matmul")(wildcard(), wildcard()) -def make_conv2d_pattern(): - return is_op("nn.conv2d")(wildcard(), wildcard()) +def make_conv2d_pattern(with_bias=False, with_act=None): + """Create a pattern for dense op followed by activations.""" + data = wildcard() + weight = wildcard() + bias = wildcard() + conv2d = is_op("nn.conv2d")(data, weight) + if with_bias: + add_or_bias_add = is_op("add") | is_op("nn.bias_add") + conv2d_out = add_or_bias_add(conv2d, bias) + else: + conv2d_out = conv2d + + if with_act is not None: + if with_act == "relu": + return is_op("nn.relu")(conv2d_out) + if with_act == "sigmoid": + return is_op("sigmoid")(conv2d_out) + + return conv2d_out def check_dtype(lhs, rhs): @@ -131,7 +148,17 @@ def partition_for_cutlass(mod): dense_bias_pat, dense_pat, ("cutlass.batch_matmul", make_batch_matmul_pattern(), check_batch_matmul), - # TODO(masahi): Add more conv2d patterns + ( + "cutlass.conv2d_bias_relu", + make_conv2d_pattern(with_bias=True, with_act="relu"), + check_conv2d, + ), + ( + "cutlass.conv2d_bias_sigmoid", + make_conv2d_pattern(with_bias=True, with_act="sigmoid"), + check_conv2d, + ), + ("cutlass.conv2d_bias", make_conv2d_pattern(with_bias=True), check_conv2d), ("cutlass.conv2d", make_conv2d_pattern(), check_conv2d), ] seq = Sequential( diff --git a/src/relay/backend/contrib/cutlass/codegen.cc b/src/relay/backend/contrib/cutlass/codegen.cc index c226da5864fc..d06ebaa896f4 100644 --- a/src/relay/backend/contrib/cutlass/codegen.cc +++ b/src/relay/backend/contrib/cutlass/codegen.cc @@ -263,6 +263,11 @@ Str2StrMap Conv2dArgs(const Map& attrs) { std::string Conv2dOp(std::string id, const Str2StrMap& attrs, const std::vector& func_args) { + bool has_bias = attrs.at("op_type") == "cutlass.conv2d_bias" || + attrs.at("op_type") == "cutlass.conv2d_bias_relu" || + attrs.at("op_type") == "cutlass.conv2d_bias_sigmoid"; + bool no_bias_scaling = attrs.at("op_type") != "cutlass.conv2d_bias_sigmoid"; + std::ostringstream conv2d_decl; CutlassPrint(conv2d_decl, "using ElementInputA = " + attrs.at("ElementInputA") + ";\n"); CutlassPrint(conv2d_decl, "using ElementInputB = " + attrs.at("ElementInputB") + ";\n"); @@ -307,10 +312,18 @@ std::string Conv2dOp(std::string id, const Str2StrMap& attrs, ICHECK(func_args.size() >= 2); CutlassPrint(conv2d_decl, "void* ptr_a = (void*)(" + func_args[0] + "->data);\n"); CutlassPrint(conv2d_decl, "void* ptr_b = (void*)(" + func_args[1] + "->data);\n"); + if (has_bias) { + ICHECK(func_args.size() >= 3); + CutlassPrint(conv2d_decl, "void* ptr_c_bias = (void*)(" + func_args[2] + "->data);\n"); + } + CutlassPrint(conv2d_decl, "void* ptr_out = (void*)(out0->data);\n"); CutlassPrint(conv2d_decl, "ElementComputeEpilogue alpha = ElementComputeEpilogue(1);\n"); - CutlassPrint(conv2d_decl, "ElementComputeEpilogue beta = ElementComputeEpilogue(0);\n"); - + if (has_bias && no_bias_scaling) { + CutlassPrint(conv2d_decl, "ElementComputeEpilogue beta = ElementComputeEpilogue(0);\n"); + } else { + CutlassPrint(conv2d_decl, "ElementComputeEpilogue beta = ElementComputeEpilogue(1);\n"); + } CutlassPrint(conv2d_decl, "using cutlass::layout::TensorNHWC;\n"); CutlassPrint(conv2d_decl, "TensorNHWC layout_A(TensorNHWC::packed(cutlass::make_Coord(N, H, W, C)));\n"); @@ -322,9 +335,19 @@ std::string Conv2dOp(std::string id, const Str2StrMap& attrs, CutlassPrint(conv2d_decl, " problem_size,\n"); CutlassPrint(conv2d_decl, " {static_cast(ptr_a), layout_A},\n"); CutlassPrint(conv2d_decl, " {static_cast(ptr_b), layout_B},\n"); + if (has_bias) { + CutlassPrint( + conv2d_decl, + " {static_cast(ptr_c_bias), cutlass::layout::TensorNHWC::Stride(0)},\n"); + } else { + CutlassPrint(conv2d_decl, " {static_cast(ptr_out),layout_C},\n"); + } CutlassPrint(conv2d_decl, " {static_cast(ptr_out),layout_C},\n"); - CutlassPrint(conv2d_decl, " {static_cast(ptr_out),layout_C},\n"); - CutlassPrint(conv2d_decl, "{alpha, beta}\n};\n"); + if (has_bias && no_bias_scaling) { + CutlassPrint(conv2d_decl, " {alpha}\n};\n"); + } else { + CutlassPrint(conv2d_decl, "{alpha, beta}\n};\n"); + } CutlassPrint(conv2d_decl, "Conv2d conv2d_op;\n"); CutlassPrint(conv2d_decl, "size_t workspace_size = conv2d_op.get_workspace_size(arguments);\n"); @@ -461,6 +484,27 @@ class CodegenCutlass : public MemoizedExprTranslator>, publi const auto* conv2d_call = GetRootCall(callee->body.as(), 0, {"nn.conv2d"}); return GenerateBody(conv2d_call, "cutlass_conv2d", GetArgumentNames(caller), Conv2dArgs(std::ref(attrs_))); + } else if (pattern_name == "cutlass.conv2d_bias") { + const CallNode* current_call = callee->body.as(); + std::string add_or_bias_add = current_call->op.as()->name; + const auto* conv2d_call = + GetRootCall(callee->body.as(), 1, {"nn.conv2d", add_or_bias_add}); + return GenerateBody(conv2d_call, "cutlass_conv2d_bias", GetArgumentNames(caller), + Conv2dArgs(std::ref(attrs_))); + } else if (pattern_name == "cutlass.conv2d_bias_relu") { + const CallNode* current_call = callee->body.as(); + std::string add_or_bias_add = current_call->args[0].as()->op.as()->name; + const auto* conv2d_call = + GetRootCall(callee->body.as(), 2, {"nn.conv2d", add_or_bias_add, "nn.relu"}); + return GenerateBody(conv2d_call, "cutlass_conv2d_bias_relu", GetArgumentNames(caller), + Conv2dArgs(std::ref(attrs_))); + } else if (pattern_name == "cutlass.conv2d_bias_sigmoid") { + const CallNode* current_call = callee->body.as(); + std::string add_or_bias_add = current_call->args[0].as()->op.as()->name; + const auto* conv2d_call = + GetRootCall(callee->body.as(), 2, {"nn.conv2d", add_or_bias_add, "sigmoid"}); + return GenerateBody(conv2d_call, "cutlass_conv2d_bias_sigmoid", GetArgumentNames(caller), + Conv2dArgs(std::ref(attrs_))); } LOG(FATAL) << "Unknown composite function: " << pattern_name; @@ -507,7 +551,9 @@ class CodegenCutlass : public MemoizedExprTranslator>, publi ret.decl = DenseOp(ext_func_id_, attribute_args, func_args); } else if (func_name == "cutlass_batch_matmul") { ret.decl = BatchMatmulOp(ext_func_id_, attribute_args, func_args); - } else if (func_name == "cutlass_conv2d") { + } else if (func_name == "cutlass_conv2d" || func_name == "cutlass_conv2d_bias" || + func_name == "cutlass_conv2d_bias_relu" || + func_name == "cutlass_conv2d_bias_sigmoid") { ret.decl = Conv2dOp(ext_func_id_, attribute_args, func_args); } diff --git a/tests/python/contrib/test_cutlass.py b/tests/python/contrib/test_cutlass.py index fee84d252081..89099c86dc58 100644 --- a/tests/python/contrib/test_cutlass.py +++ b/tests/python/contrib/test_cutlass.py @@ -114,18 +114,30 @@ def get_conv2d_nchw(d_shape, w_shape, padding, out_dtype="float16"): data = relay.var("data", shape=d_shape, dtype="float16") weight = relay.var("weight", shape=w_shape, dtype="float16") out_channel = w_shape[0] - return tvm.IRModule.from_expr( - relay.nn.conv2d( - data=data, - weight=weight, - kernel_size=w_shape[2:], - channels=out_channel, - padding=padding, - out_dtype=out_dtype, - ) + return relay.nn.conv2d( + data=data, + weight=weight, + kernel_size=w_shape[2:], + channels=out_channel, + padding=padding, + out_dtype=out_dtype, ) +def get_conv2d_nchw_bias(d_shape, w_shape, padding, out_dtype="float16"): + conv2d = get_conv2d_nchw(d_shape, w_shape, padding, out_dtype=out_dtype) + bias = relay.var("bias", shape=(w_shape[0],), dtype=out_dtype) + return relay.nn.bias_add(conv2d, bias) + + +def get_conv2d_nchw_bias_relu(d_shape, w_shape, padding, out_dtype="float16"): + return relay.nn.relu(get_conv2d_nchw_bias(d_shape, w_shape, padding, out_dtype=out_dtype)) + + +def get_conv2d_nchw_bias_sigmoid(d_shape, w_shape, padding, out_dtype="float16"): + return relay.sigmoid(get_conv2d_nchw_bias(d_shape, w_shape, padding, out_dtype=out_dtype)) + + def profile_and_build(mod, params, sm, tmp_dir="./tmp", lib_path="compile.so"): mod = partition_for_cutlass(mod) mod, num_cutlass_partition = tune_cutlass_kernels( @@ -314,8 +326,8 @@ def convert_conv2d_layout(mod, desired_layouts): def verify_conv2d( - mod_nchw, # can be dynamic batch - mod_ref, # always static batch + expr_nchw, # can be dynamic batch + expr_ref, # always static batch d_shape, w_shape, sm=80, @@ -327,10 +339,17 @@ def verify_conv2d( if not has_cutlass(): return + mod_nchw = tvm.IRModule.from_expr(expr_nchw) + mod_ref = tvm.IRModule.from_expr(expr_ref) + + typ = relay.transform.InferType()(mod_nchw)["main"].body.checked_type + out_dtype = typ.dtype + np_data = np.random.uniform(-1, 1, d_shape).astype("float16") np_weight = np.random.uniform(-1, 1, w_shape).astype("float16") + np_bias = np.random.uniform(-1, 1, (w_shape[0],)).astype(out_dtype) - params = {"weight": np_weight} + params = {"weight": np_weight, "bias": np_bias} typ = relay.transform.InferType()(mod_nchw)["main"].body.checked_type use_vm = any(isinstance(s, tvm.tir.Any) for s in typ.shape) @@ -373,10 +392,10 @@ def verify_conv2d( def test_conv2d(): + padding = (1, 1) for IC in [3, 16]: d_shape = (16, IC, 32, 32) w_shape = (32, IC, 3, 3) - padding = (1, 1) mod_nchw = get_conv2d_nchw(d_shape, w_shape, padding) verify_conv2d( @@ -404,5 +423,26 @@ def test_conv2d(): ) +def test_conv2d_fusion(): + d_shape = (16, 16, 32, 32) + w_shape = (32, 16, 3, 3) + padding = (1, 1) + + mod_nchw = get_conv2d_nchw_bias(d_shape, w_shape, padding) + verify_conv2d( + mod_nchw, mod_nchw, d_shape, w_shape, sm=80, atol=1e-5, rtol=1e-5, run_benchmark=False + ) + + mod_nchw = get_conv2d_nchw_bias_relu(d_shape, w_shape, padding) + verify_conv2d( + mod_nchw, mod_nchw, d_shape, w_shape, sm=80, atol=1e-5, rtol=1e-5, run_benchmark=False + ) + + mod_nchw = get_conv2d_nchw_bias_sigmoid(d_shape, w_shape, padding, out_dtype="float32") + verify_conv2d( + mod_nchw, mod_nchw, d_shape, w_shape, sm=80, atol=1e-5, rtol=1e-5, run_benchmark=False + ) + + if __name__ == "__main__": pytest.main([__file__]) From 48cac9496d6ecc76a746506f3caa2c94c551be2d Mon Sep 17 00:00:00 2001 From: Masahiro Masuda Date: Mon, 13 Dec 2021 21:28:48 +0900 Subject: [PATCH 2/2] support batch norm fusion --- python/tvm/relay/op/contrib/cutlass.py | 22 +++++++++++++++++++--- 1 file changed, 19 insertions(+), 3 deletions(-) diff --git a/python/tvm/relay/op/contrib/cutlass.py b/python/tvm/relay/op/contrib/cutlass.py index c706769b2d3d..8fdd90ea109a 100644 --- a/python/tvm/relay/op/contrib/cutlass.py +++ b/python/tvm/relay/op/contrib/cutlass.py @@ -16,8 +16,9 @@ # under the License. # pylint: disable=invalid-name """Patterns supported CUTLASS.""" -from tvm.ir.transform import Sequential +from tvm.ir.transform import Sequential, PassContext from tvm.relay import transform +from tvm.relay.build_module import bind_params_by_name from ...dataflow_pattern import wildcard, is_op, is_constant @@ -126,7 +127,7 @@ def check_conv2d(call): return not is_depthwise_conv2d(IC, OC, conv2d.attrs.groups) -def partition_for_cutlass(mod): +def partition_for_cutlass(mod, params=None): """Partition the input module into CUTLASS-supported subgraphs.""" dense_pat = ("cutlass.dense", make_gemm_pattern(False, None), check_gemm) dense_bias_pat = ("cutlass.dense_bias", make_gemm_pattern(True, None), check_gemm) @@ -161,12 +162,27 @@ def partition_for_cutlass(mod): ("cutlass.conv2d_bias", make_conv2d_pattern(with_bias=True), check_conv2d), ("cutlass.conv2d", make_conv2d_pattern(), check_conv2d), ] + + if params is not None: + mod["main"] = bind_params_by_name(mod["main"], params) + remove_bn_pass = Sequential( + [ + transform.InferType(), + transform.SimplifyInference(), + transform.FoldConstant(), + transform.FoldScaleAxis(), + ] + ) + with PassContext(opt_level=3): + mod = remove_bn_pass(mod) + seq = Sequential( [ transform.InferType(), transform.MergeComposite(cutlass_patterns), transform.AnnotateTarget(["cutlass"]), - transform.PartitionGraph(), + transform.PartitionGraph(bind_constants=False), ] ) + return seq(mod)