From 4cd55c33d8b732adf59d1ca41a28f17b7b3f1b0f Mon Sep 17 00:00:00 2001 From: ganler Date: Wed, 24 Nov 2021 16:10:05 -0600 Subject: [PATCH 1/7] fix: integer mismatch in type inference by lifting constant int32 to int64 symbols --- src/tir/ir/data_layout.cc | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/tir/ir/data_layout.cc b/src/tir/ir/data_layout.cc index da3496dba407..83ca2a90ff37 100644 --- a/src/tir/ir/data_layout.cc +++ b/src/tir/ir/data_layout.cc @@ -208,7 +208,7 @@ inline bool GetStoreRule(Array* rule, const Layout& src_layout, for (size_t i = 0; i < dst_layout.ndim(); ++i) { const auto& store_axis = dst_layout[i]; const IterVar& store_axis_impl = dst_layout->axes[i]; - PrimExpr store(0); + PrimExpr store(IntImm(DataType::Int(64), 0)); for (size_t j = 0; j < src_layout.ndim(); ++j) { const auto& orig_axis = src_layout[j]; @@ -218,7 +218,7 @@ inline bool GetStoreRule(Array* rule, const Layout& src_layout, PrimExpr orig_var = orig_axis_impl->var; const int32_t factor = src_layout.FactorOf(orig_axis); if (factor > 0) { - orig_var = orig_var * PrimExpr(factor); + orig_var = orig_var * PrimExpr(IntImm(DataType::Int(64), factor)); } store = store + orig_var; } else { @@ -304,7 +304,7 @@ inline Array TransformShape(const Array& src_shape, << ", get " << orig_shape; } } - bind_map[orig_axis->var.get()] = PrimExpr(0); + bind_map[orig_axis->var.get()] = PrimExpr(IntImm(DataType::Int(64), 0)); } else { bind_map[orig_axis->var.get()] = orig_shape; } From 2402511ed203b756ed2a61e7bd5d54858538b210 Mon Sep 17 00:00:00 2001 From: ganler Date: Fri, 26 Nov 2021 15:07:35 -0600 Subject: [PATCH 2/7] refine Ramp node checking to pass tests --- src/tir/ir/expr.cc | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/tir/ir/expr.cc b/src/tir/ir/expr.cc index 1d7c959d990d..a306d5eaa119 100644 --- a/src/tir/ir/expr.cc +++ b/src/tir/ir/expr.cc @@ -707,7 +707,8 @@ Ramp::Ramp(PrimExpr base, PrimExpr stride, int lanes, Span span) { ICHECK(base.dtype().is_scalar()); ICHECK(stride.dtype().is_scalar()); ICHECK_GT(lanes, 1); - ICHECK_EQ(stride.dtype(), base.dtype()); + ICHECK(base.dtype().is_int()); + ICHECK(stride.dtype().is_int()); ObjectPtr node = make_object(); node->dtype = base.dtype().with_lanes(lanes); From 7e6a7da93c547a1c531ea6e050c18bf3fb4477c8 Mon Sep 17 00:00:00 2001 From: ganler Date: Fri, 26 Nov 2021 17:09:28 -0600 Subject: [PATCH 3/7] add test of integer compatibility testing for layout transform --- tests/python/relay/test_type_solver.py | 32 ++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) diff --git a/tests/python/relay/test_type_solver.py b/tests/python/relay/test_type_solver.py index 88bdd1628920..81b4f22f832c 100644 --- a/tests/python/relay/test_type_solver.py +++ b/tests/python/relay/test_type_solver.py @@ -16,7 +16,10 @@ # under the License. import tvm from tvm import relay +from tvm.relay import testing + import pytest +import numpy as np def make_rel(name, args, num_inputs=None, attrs=None): @@ -338,6 +341,34 @@ def test_incompatible_quantified_func_unification(): solver.Unify(ft1, ft2) +def test_integer_compatibility_in_layout_transform(): + x = relay.var("data", shape=(2, 3, 48, 48), dtype="float32") + conv_out = relay.nn.conv2d( + x, + relay.var("weight", shape=(1, 3, 1, 1), dtype="float32"), + strides=[47, 47], + padding=[0, 0, 0, 0], + channels=1, + kernel_size=[1, 1], + ) + bias_out = relay.nn.bias_add(conv_out, relay.var("bias")) + broadcast_out = relay.op.broadcast_to(bias_out, relay.const([2, 1, 2, 2], dtype="int64")) + y = relay.add(bias_out, broadcast_out) + + mod, params = testing.create_workload(y) + with tvm.transform.PassContext(opt_level=3): + executor = relay.build_module.create_executor( + "graph", + mod, + tvm.cpu(), + "llvm", + params={ + "weight": np.zeros((1, 3, 1, 1), dtype="float32"), + "bias": np.zeros((1), dtype="float32"), + }, + ).evaluate() + + if __name__ == "__main__": test_bcast() test_backward_solving() @@ -357,3 +388,4 @@ def test_incompatible_quantified_func_unification(): test_incompatible_typecall_var_unification() test_incompatible_typecall_args_unification() test_incompatible_quantified_func_unification() + test_integer_compatibility_in_layout_transform() From 24797aee0c4cbfe6883dcfd795ddd79de7c2e864 Mon Sep 17 00:00:00 2001 From: ganler Date: Fri, 26 Nov 2021 20:40:56 -0600 Subject: [PATCH 4/7] refine: use a compatible fix --- src/tir/ir/data_layout.cc | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/src/tir/ir/data_layout.cc b/src/tir/ir/data_layout.cc index 83ca2a90ff37..8dea3435ccab 100644 --- a/src/tir/ir/data_layout.cc +++ b/src/tir/ir/data_layout.cc @@ -208,7 +208,7 @@ inline bool GetStoreRule(Array* rule, const Layout& src_layout, for (size_t i = 0; i < dst_layout.ndim(); ++i) { const auto& store_axis = dst_layout[i]; const IterVar& store_axis_impl = dst_layout->axes[i]; - PrimExpr store(IntImm(DataType::Int(64), 0)); + PrimExpr store(0); for (size_t j = 0; j < src_layout.ndim(); ++j) { const auto& orig_axis = src_layout[j]; @@ -218,7 +218,7 @@ inline bool GetStoreRule(Array* rule, const Layout& src_layout, PrimExpr orig_var = orig_axis_impl->var; const int32_t factor = src_layout.FactorOf(orig_axis); if (factor > 0) { - orig_var = orig_var * PrimExpr(IntImm(DataType::Int(64), factor)); + orig_var = orig_var * factor; } store = store + orig_var; } else { @@ -304,9 +304,11 @@ inline Array TransformShape(const Array& src_shape, << ", get " << orig_shape; } } - bind_map[orig_axis->var.get()] = PrimExpr(IntImm(DataType::Int(64), 0)); + bind_map[orig_axis->var.get()] = IntImm(orig_axis->var->dtype, 0); } else { - bind_map[orig_axis->var.get()] = orig_shape; + bind_map[orig_axis->var.get()] = orig_axis->var->dtype == orig_shape->dtype + ? orig_shape + : cast(orig_axis->var->dtype, orig_shape); } } // infer the target shape, From 71fc565843c8baccbc8978536f2421cb20d81014 Mon Sep 17 00:00:00 2001 From: ganler Date: Mon, 13 Dec 2021 01:01:34 -0600 Subject: [PATCH 5/7] minimize change --- src/tir/ir/expr.cc | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/tir/ir/expr.cc b/src/tir/ir/expr.cc index a306d5eaa119..1d7c959d990d 100644 --- a/src/tir/ir/expr.cc +++ b/src/tir/ir/expr.cc @@ -707,8 +707,7 @@ Ramp::Ramp(PrimExpr base, PrimExpr stride, int lanes, Span span) { ICHECK(base.dtype().is_scalar()); ICHECK(stride.dtype().is_scalar()); ICHECK_GT(lanes, 1); - ICHECK(base.dtype().is_int()); - ICHECK(stride.dtype().is_int()); + ICHECK_EQ(stride.dtype(), base.dtype()); ObjectPtr node = make_object(); node->dtype = base.dtype().with_lanes(lanes); From 43cd93815398041c5922c13a8cae2a86dfa7b066 Mon Sep 17 00:00:00 2001 From: ganler Date: Mon, 13 Dec 2021 20:11:58 -0600 Subject: [PATCH 6/7] [EMPTY] trigger bad CI From d9d69a67480441fe95c5ad1abc18c529b99bece7 Mon Sep 17 00:00:00 2001 From: ganler Date: Sun, 2 Jan 2022 20:45:35 -0600 Subject: [PATCH 7/7] refact: minimize int mismatch test case --- tests/python/relay/test_type_solver.py | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/tests/python/relay/test_type_solver.py b/tests/python/relay/test_type_solver.py index 81b4f22f832c..c1dc5c03a420 100644 --- a/tests/python/relay/test_type_solver.py +++ b/tests/python/relay/test_type_solver.py @@ -347,7 +347,6 @@ def test_integer_compatibility_in_layout_transform(): x, relay.var("weight", shape=(1, 3, 1, 1), dtype="float32"), strides=[47, 47], - padding=[0, 0, 0, 0], channels=1, kernel_size=[1, 1], ) @@ -355,18 +354,11 @@ def test_integer_compatibility_in_layout_transform(): broadcast_out = relay.op.broadcast_to(bias_out, relay.const([2, 1, 2, 2], dtype="int64")) y = relay.add(bias_out, broadcast_out) - mod, params = testing.create_workload(y) + mod, _ = testing.create_workload(y) with tvm.transform.PassContext(opt_level=3): - executor = relay.build_module.create_executor( - "graph", - mod, - tvm.cpu(), - "llvm", - params={ - "weight": np.zeros((1, 3, 1, 1), dtype="float32"), - "bias": np.zeros((1), dtype="float32"), - }, - ).evaluate() + with tvm.target.Target("llvm"): + mod = relay.transform.CanonicalizeOps()(mod) + mod = relay.transform.AlterOpLayout()(mod) if __name__ == "__main__":