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

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


The following commit(s) were added to refs/heads/main by this push:
     new a6f6f11000 [TensorIR][Bugfix] `reindex_cache_write` do not mutate init 
statement (#14626)
a6f6f11000 is described below

commit a6f6f11000860438e2ea7d80a2c017ef78d646df
Author: Zihao Ye <[email protected]>
AuthorDate: Sat Apr 15 08:03:30 2023 -0700

    [TensorIR][Bugfix] `reindex_cache_write` do not mutate init statement 
(#14626)
    
    # The Bug
    When applying `reindex_cache_write` to the write buffer of a reduction 
block, the init statement would not be mutated accordingly:
    
    ```python
    # original program
    @T.prim_func
    def reduce(A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 
128), "float32")):
        B = T.alloc_buffer((128, 128, 128), dtype="float32")
        for i, j, k in T.grid(128, 128, 128):
            for l in range(128):
                with T.block("B"):
                    vi, vj, vk, vl = T.axis.remap("SSSR", [i, j, k, l])
                    with T.init():
                        B[vi, vj, vk] = T.float32(0)
                    B[vi, vj, vk] = B[vi, vj, vk] + A[vi, vj, vk, vl]
            with T.block("C"):
                vi, vj, vk = T.axis.remap("SSR", [i, j, k])
                with T.init():
                    C[vi, vj] = T.float32(0)
                C[vi, vj] = C[vi, vj] + B[vi, vj, vk]
    
    # schedule
    sch = tir.Schedule(reduce, debug_mask="all")
    sch.reindex_cache_write("B", 0, "shared", lambda i, j, k, l: (j, i, k))
    
    # after schedule
    @T.prim_func
    def reduce_after_reindex_cache_write(
        A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), 
"float32")
    ):
        B = T.alloc_buffer((128, 128, 128))
        B_shared = T.alloc_buffer((128, 128, 128), scope="shared")
        for i, j, k in T.grid(128, 128, 128):
            for l in range(128):
                with T.block("B"):
                    vi, vj, vk, vl = T.axis.remap("SSSR", [i, j, k, l])
                    T.reads(A[vi, vj, vk, vl])
                    T.writes(B_shared[vj, vi, vk])
                    with T.init():
                        B[vj, vi, vk] = T.float32(0)
                    B_shared[vj, vi, vk] = B_shared[vj, vi, vk] + A[vi, vj, vk, 
vl]
            with T.block("B_shared"):
                vi, vj, vk = T.axis.remap("SSS", [i, j, k])
                T.reads(B_shared[vj, vi, vk])
                T.writes(B[vi, vj, vk])
                B[vi, vj, vk] = B_shared[vj, vi, vk]
            with T.block("C"):
                vi, vj, vk = T.axis.remap("SSR", [i, j, k])
                T.reads(B[vi, vj, vk])
                T.writes(C[vi, vj])
                with T.init():
                    C[vi, vj] = T.float32(0)
                C[vi, vj] = C[vi, vj] + B[vi, vj, vk]
    ```
    
    The init statement inside block "B" should be transformed to `B[vj, vi, vk] 
= T.float32(0)`
    
    # The Fix
    In our previous implementation, we mistakenly specify the consumer block to 
be the block itself, which is not necessary and would cause the later 
`ReindexCacheWriteRewriter` to skip rewriting the init statement.
---
 src/tir/schedule/primitive/cache_read_write.cc     |  2 -
 .../unittest/test_tir_schedule_cache_read_write.py | 91 ++++++++++++++++++++++
 2 files changed, 91 insertions(+), 2 deletions(-)

diff --git a/src/tir/schedule/primitive/cache_read_write.cc 
b/src/tir/schedule/primitive/cache_read_write.cc
index c2fc7ac24a..cf139c7df7 100644
--- a/src/tir/schedule/primitive/cache_read_write.cc
+++ b/src/tir/schedule/primitive/cache_read_write.cc
@@ -1842,8 +1842,6 @@ StmtSRef ReindexCacheWrite(ScheduleState self, const 
StmtSRef& block_sref, int w
   // Step 2. Creating CacheStageInfo
   ReindexCacheStageInfo info;
   info.write_buffer = write_buffer;
-  LOG(INFO) << block->name_hint;
-  info.consumer_blocks.insert(block_sref);
 
   // Step 3. Check the only writer block.
   ICHECK_EQ(block_sref.get(), GetOnlyWriteBlock(self, scope_sref, 
write_buffer).get());
diff --git a/tests/python/unittest/test_tir_schedule_cache_read_write.py 
b/tests/python/unittest/test_tir_schedule_cache_read_write.py
index cf75768ec0..454557a2bd 100644
--- a/tests/python/unittest/test_tir_schedule_cache_read_write.py
+++ b/tests/python/unittest/test_tir_schedule_cache_read_write.py
@@ -111,6 +111,88 @@ def elementwise_reindex_cache_write(
             C[vi, vj] = B[vi, vj] + T.float32(1)
 
 
[email protected]_func
+def reduce(A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 
128), "float32")):
+    B = T.alloc_buffer((128, 128, 128), dtype="float32")
+    for i, j, k in T.grid(128, 128, 128):
+        for l in range(128):
+            with T.block("B"):
+                vi, vj, vk, vl = T.axis.remap("SSSR", [i, j, k, l])
+                with T.init():
+                    B[vi, vj, vk] = T.float32(0)
+                B[vi, vj, vk] = B[vi, vj, vk] + A[vi, vj, vk, vl]
+        with T.block("C"):
+            vi, vj, vk = T.axis.remap("SSR", [i, j, k])
+            with T.init():
+                C[vi, vj] = T.float32(0)
+            C[vi, vj] = C[vi, vj] + B[vi, vj, vk]
+
+
[email protected]_func
+def reduce_reindex_cache_write_0(
+    A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), 
"float32")
+):
+    B = T.alloc_buffer((128, 128, 128))
+    B_shared = T.alloc_buffer((128, 128, 128), scope="shared")
+    for i, j, k in T.grid(128, 128, 128):
+        for l in range(128):
+            with T.block("B"):
+                vi, vj, vk, vl = T.axis.remap("SSSR", [i, j, k, l])
+                T.reads(A[vi, vj, vk, vl])
+                T.writes(B_shared[vj, vi, vk])
+                with T.init():
+                    B_shared[vj, vi, vk] = T.float32(0)
+                B_shared[vj, vi, vk] = B_shared[vj, vi, vk] + A[vi, vj, vk, vl]
+        with T.block("B_shared"):
+            vi, vj, vk = T.axis.remap("SSS", [i, j, k])
+            T.reads(B_shared[vj, vi, vk])
+            T.writes(B[vi, vj, vk])
+            B[vi, vj, vk] = B_shared[vj, vi, vk]
+        with T.block("C"):
+            vi, vj, vk = T.axis.remap("SSR", [i, j, k])
+            T.reads(B[vi, vj, vk])
+            T.writes(C[vi, vj])
+            with T.init():
+                C[vi, vj] = T.float32(0)
+            C[vi, vj] = C[vi, vj] + B[vi, vj, vk]
+
+
[email protected]_func
+def reduce_reindex_cache_write_1(
+    A: T.Buffer((128, 128, 128, 128), "float32"), C: T.Buffer((128, 128), 
"float32")
+):
+    B = T.alloc_buffer((128, 128, 128))
+    B_shared = T.alloc_buffer((128, 128, 128), scope="shared")
+    C_shared = T.alloc_buffer((128, 128), scope="shared")
+    for i, j, k in T.grid(128, 128, 128):
+        for l in range(128):
+            with T.block("B"):
+                vi, vj, vk, vl = T.axis.remap("SSSR", [i, j, k, l])
+                T.reads(A[vi, vj, vk, vl])
+                T.writes(B_shared[vj, vi, vk])
+                with T.init():
+                    B_shared[vj, vi, vk] = T.float32(0)
+                B_shared[vj, vi, vk] = B_shared[vj, vi, vk] + A[vi, vj, vk, vl]
+        with T.block("B_shared"):
+            vi, vj, vk = T.axis.remap("SSS", [i, j, k])
+            T.reads(B_shared[vj, vi, vk])
+            T.writes(B[vi, vj, vk])
+            B[vi, vj, vk] = B_shared[vj, vi, vk]
+        with T.block("C"):
+            vi, vj, vk = T.axis.remap("SSR", [i, j, k])
+            T.reads(B[vi, vj, vk])
+            T.writes(C_shared[vj, vi])
+            with T.init():
+                C_shared[vj, vi] = T.float32(0)
+            C_shared[vj, vi] = C_shared[vj, vi] + B[vi, vj, vk]
+    for i, j in T.grid(128, 128):
+        with T.block("C_shared"):
+            vi, vj = T.axis.remap("SS", [i, j])
+            T.reads(C_shared[vj, vi])
+            T.writes(C[vi, vj])
+            C[vi, vj] = C_shared[vj, vi]
+
+
 @T.prim_func
 def func_nested_seq(b: T.handle, c: T.handle) -> None:
     A = T.alloc_buffer((128, 128))
@@ -1459,6 +1541,15 @@ def test_reindex_cache_write():
     verify_trace_roundtrip(sch=sch, mod=elementwise)
 
 
+def test_reindex_cache_write_reduce():
+    sch = tir.Schedule(reduce, debug_mask="all")
+    sch.reindex_cache_write("B", 0, "shared", lambda i, j, k, l: (j, i, k))
+    tvm.ir.assert_structural_equal(reduce_reindex_cache_write_0, 
sch.mod["main"])
+    sch.reindex_cache_write("C", 0, "shared", lambda i, j, k: [j, i])
+    tvm.ir.assert_structural_equal(reduce_reindex_cache_write_1, 
sch.mod["main"])
+    verify_trace_roundtrip(sch=sch, mod=reduce)
+
+
 def test_reindex_cache_write_fail_not_match():
     sch = tir.Schedule(elementwise, debug_mask="all")
     with pytest.raises(tvm.tir.ScheduleError):

Reply via email to