This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/unity by this push:
     new f796b69a3b [Unity] Fix LazyTransformParams use-def analysis and 
binding emission (#14974)
f796b69a3b is described below

commit f796b69a3baac50b21c24012993a03e5a117f25f
Author: Ruihang Lai <[email protected]>
AuthorDate: Tue May 30 11:37:50 2023 -0700

    [Unity] Fix LazyTransformParams use-def analysis and binding emission 
(#14974)
    
    Previously the pass LazyTransformParams did not take the case where an
    output variable is used by other bindings inside the function. It
    generates wrong Relax function in this case.
    
    This PR fixes this issue. The new test case is an example to elaborate
    the issue. This PR introduces a sanity assertion check to ensure we
    handle the case properly.
    
    This PR also enhances variable naming. For binding variables of
    `set_item` and `kill_object`, we now use an underscore ("`_`") as its
    name, compared with the previous `lv` name.
---
 .../tvm/relax/transform/lazy_transform_params.py   | 72 ++++++++++++----------
 .../relax/test_transform_lazy_transform_params.py  | 66 +++++++++++++++++---
 2 files changed, 96 insertions(+), 42 deletions(-)

diff --git a/python/tvm/relax/transform/lazy_transform_params.py 
b/python/tvm/relax/transform/lazy_transform_params.py
index 9ce7bd0038..fe532033d9 100644
--- a/python/tvm/relax/transform/lazy_transform_params.py
+++ b/python/tvm/relax/transform/lazy_transform_params.py
@@ -79,20 +79,20 @@ class LivenessAnalysis(PyExprVisitor):
         The set of vars that are bound to v = params[i]
     """
 
-    def __init__(self, out_tuple_var: relax.Var, input_params: set) -> None:
+    def __init__(self, out_tuple_var: relax.Var) -> None:
         self.last_appear_in_var_binding = None
         self.out_tuple_var = out_tuple_var
-        self.input_params = input_params
         self.var_liveness_end = {}
+        self.ended_vars = set()
 
     def visit_binding_block_(self, block: relax.BindingBlock) -> None:
         for binding in reversed(block.bindings):
             self.visit_binding(binding)
 
     def visit_var_(self, op: relax.Var) -> None:
-        if op in self.input_params:
+        if op not in self.ended_vars:
             self.last_appear_in_var_binding.append(op)
-            self.input_params.remove(op)
+            self.ended_vars.add(op)
 
     def visit_var_binding_(self, binding: relax.VarBinding) -> None:
         if self.out_tuple_var == binding.var:
@@ -100,8 +100,9 @@ class LivenessAnalysis(PyExprVisitor):
         self.last_appear_in_var_binding = []
         super().visit_var_binding_(binding)
         # param[i] is in output
-        if binding.var in self.input_params:
+        if binding.var not in self.ended_vars:
             self.last_appear_in_var_binding.append(binding.var)
+            self.ended_vars.add(binding.var)
         self.var_liveness_end[binding.var] = self.last_appear_in_var_binding
 
 
@@ -119,14 +120,13 @@ class LazyTransformParamsMutator(PyExprMutator):
     def __init__(self, mod: IRModule = None) -> None:
         super().__init__(mod)
         self.mod = mod
-        self.get_item = None
-        self.set_item = None
         # the only input param, which should be a Tuple
         self.input_tuple_param = None
-        # map from out var to index
+        self.input_params_set = None
         self.out_tuple_map = None
         self.out_tuple_var = None
         self.memory_free_insertion = None
+        self.killed_vars = set()
 
     def transform(self, func: relax.Function) -> relax.Function:
         self.input_tuple_param = func.params[0]
@@ -137,9 +137,9 @@ class LazyTransformParamsMutator(PyExprMutator):
         forward_collector.visit_expr(func)
         self.out_tuple_map = forward_collector.out_tuple_map
         # input_params_set is the set of binding var for var = params[i]
-        input_params_set = set(forward_collector.var_tuple_get_item)
+        self.input_params_set = set(forward_collector.var_tuple_get_item)
         # Step 2. liveness analysis and get where to insert kill_object 
instruction
-        liveness = LivenessAnalysis(self.out_tuple_var, input_params_set)
+        liveness = LivenessAnalysis(self.out_tuple_var)
         liveness.visit_expr(func)
         self.memory_free_insertion = liveness.var_liveness_end
         # Step 3. rewrite get item and set item
@@ -159,31 +159,39 @@ class LazyTransformParamsMutator(PyExprMutator):
         else:
             return tuple_get_item
 
+    def visit_var_(self, var: relax.Var) -> None:
+        assert var not in self.killed_vars
+        return super().visit_var_(var)
+
     def visit_var_binding_(self, binding: relax.VarBinding) -> None:
-        if binding.var in self.out_tuple_map:
-            index = self.out_tuple_map[binding.var]
-            value = self.visit_expr(binding.value)
-            var_before_setitem = self.builder_.emit(value)
-            # rewrite set item
-            new_var = self.builder_.emit(
-                relax.Call(
-                    relax.ExternFunc("set_item"),
-                    [index, var_before_setitem],
-                    None,
-                    [relax.ObjectStructInfo()],
-                )
-            )
-            self.set_var_remap(binding.var.vid, new_var)
-        else:
-            super().visit_var_binding_(binding)
+        if binding.var == self.out_tuple_var:
+            # The function after rewriting returns a empty tuple.
+            func_output = self.builder_.emit(relax.Tuple([]))
+            self.set_var_remap(binding.var.vid, func_output)
+            return
+
+        super().visit_var_binding_(binding)
+
         if binding.var in self.memory_free_insertion:
             for var in self.memory_free_insertion[binding.var]:
-                # handle param[i] in output
-                if var == binding.var:
-                    assert binding.var in self.out_tuple_map
-                    
self.builder_.emit(relax.op.vm.kill_object(var_before_setitem))
-                else:
-                    
self.builder_.emit(relax.op.vm.kill_object(self.get_var_remap(var.vid)))
+                if var in self.out_tuple_map:
+                    self.killed_vars.add(var)
+                    index = self.out_tuple_map[var]
+                    # rewrite set item
+                    self.builder_.emit(
+                        relax.Call(
+                            relax.ExternFunc("set_item"),
+                            [index, super().visit_var_(var)],
+                            None,
+                            [relax.ObjectStructInfo()],
+                        ),
+                        name_hint="_",
+                    )
+
+                if var in self.input_params_set:
+                    self.builder_.emit(
+                        relax.op.vm.kill_object(super().visit_var_(var)), 
name_hint="_"
+                    )
 
 
 @tvm.transform.module_pass(opt_level=0, name="LazyTransformParams")
diff --git a/tests/python/relax/test_transform_lazy_transform_params.py 
b/tests/python/relax/test_transform_lazy_transform_params.py
index 3de4a1ff0a..e3e454f31f 100644
--- a/tests/python/relax/test_transform_lazy_transform_params.py
+++ b/tests/python/relax/test_transform_lazy_transform_params.py
@@ -16,7 +16,6 @@
 # under the License.
 import tvm
 import tvm.testing
-from tvm import relax
 from tvm.script import relax as R, tir as T
 from tvm.script import ir as I
 from tvm.relax.transform import LazyTransformParams
@@ -75,26 +74,73 @@ def test_lazy_transform_params():
                     out[o, i, h, w] = w1[i, o, h, w]
 
         @R.function
-        def main_transform_params() -> R.Tuple(R.Object, R.Object):
+        def main_transform_params() -> R.Tuple:
             R.func_attr({"relax.force_pure": True})
             cls = Expected
             lv: R.Object = R.call_packed("get_item", R.prim_value(1), 
sinfo_args=(R.Object,))
-            lv1: R.Object = R.call_packed("set_item", R.prim_value(0), lv, 
sinfo_args=(R.Object,))
-            lv2: R.Tuple = R.vm.kill_object(lv)
-            lv1_1: R.Object = R.call_packed("get_item", R.prim_value(0), 
sinfo_args=(R.Object,))
-            lv3 = R.call_tir(
+            _: R.Object = R.call_packed("set_item", R.prim_value(0), lv, 
sinfo_args=(R.Object,))
+            _1: R.Tuple = R.vm.kill_object(lv)
+            lv1: R.Object = R.call_packed("get_item", R.prim_value(0), 
sinfo_args=(R.Object,))
+            lv2 = R.call_tir(
                 cls.transform_layout_IOHW_to_OIHW,
-                (lv1_1,),
+                (lv1,),
                 out_sinfo=R.Tensor((16, 3, 3, 3), dtype="float32"),
             )
-            lv4: R.Object = R.call_packed("set_item", R.prim_value(1), lv3, 
sinfo_args=(R.Object,))
-            lv5: R.Tuple = R.vm.kill_object(lv1_1)
-            gv: R.Tuple(R.Object, R.Object) = (lv1, lv4)
+            _2: R.Tuple = R.vm.kill_object(lv1)
+            _3: R.Object = R.call_packed("set_item", R.prim_value(1), lv2, 
sinfo_args=(R.Object,))
+            gv: R.Tuple = R.tuple()
             return gv
 
     after = LazyTransformParams()(Before)
     tvm.ir.assert_structural_equal(after, Expected, map_free_vars=True)
 
 
+def test_output_with_use_site():
+    @I.ir_module
+    class Module:
+        @T.prim_func
+        def copy(x: T.Buffer((), "float32"), y: T.Buffer((), "float32")):
+            with T.block("block"):
+                y[()] = x[()]
+
+        @R.function
+        def main_transform_params(
+            params: R.Tuple(R.Tensor((), dtype="float32"))
+        ) -> R.Tuple(R.Tensor((), dtype="float32"), R.Tensor((), 
dtype="float32")):
+            # we expect ToNonDataflow and RemovePurityTracking to be invoked 
first
+            R.func_attr({"relax.force_pure": True})
+            cls = Module
+            x: R.Tensor((), dtype="float32") = params[0]
+            y = R.call_tir(cls.copy, (x,), out_sinfo=R.Tensor((), 
dtype="float32"))
+            z = R.call_tir(cls.copy, (y,), out_sinfo=R.Tensor((), 
dtype="float32"))
+            gv: R.Tuple(R.Tensor((), dtype="float32"), R.Tensor((), 
dtype="float32")) = (y, z)
+            return gv
+
+    @I.ir_module
+    class Expected:
+        @T.prim_func
+        def copy(x: T.Buffer((), "float32"), y: T.Buffer((), "float32")):
+            with T.block("block"):
+                T.reads(x[()])
+                T.writes(y[()])
+                y[()] = x[()]
+
+        @R.function
+        def main_transform_params() -> R.Tuple:
+            R.func_attr({"relax.force_pure": True})
+            cls = Expected
+            x: R.Object = R.call_packed("get_item", R.prim_value(0), 
sinfo_args=(R.Object,))
+            y = R.call_tir(cls.copy, (x,), out_sinfo=R.Tensor((), 
dtype="float32"))
+            _: R.Tuple = R.vm.kill_object(x)
+            z = R.call_tir(cls.copy, (y,), out_sinfo=R.Tensor((), 
dtype="float32"))
+            _1: R.Object = R.call_packed("set_item", R.prim_value(0), y, 
sinfo_args=(R.Object,))
+            _2: R.Object = R.call_packed("set_item", R.prim_value(1), z, 
sinfo_args=(R.Object,))
+            gv: R.Tuple = R.tuple()
+            return gv
+
+    after = LazyTransformParams()(Module)
+    tvm.ir.assert_structural_equal(after, Expected, map_free_vars=True)
+
+
 if __name__ == "__main__":
     tvm.testing.main()

Reply via email to