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):