This is an automated email from the ASF dual-hosted git repository.
akaashrp 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 e269315c90 [Fix][S-TIR] Add a block-aware ForceNarrowIndexToInt32
(#20409)
e269315c90 is described below
commit e269315c90e3a061c9e1c77b370ce883b1b223f4
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Tue Sep 22 16:28:37 2026 -0700
[Fix][S-TIR] Add a block-aware ForceNarrowIndexToInt32 (#20409)
---
include/tvm/s_tir/transform.h | 13 ++
include/tvm/tirx/transform.h | 3 +
python/tvm/s_tir/transform/transform.py | 19 +++
python/tvm/tirx/transform/transform.py | 3 +
src/s_tir/transform/force_narrow_index_to_i32.cc | 78 +++++++++
src/tirx/transform/force_narrow_index_to_i32.cc | 64 ++------
...index_to_i32.cc => force_narrow_index_to_i32.h} | 99 ++++++------
.../relax/test_backend_dispatch_sort_scan.py | 4 +-
...st_s_tir_transform_force_narrow_index_to_i32.py | 175 +++++++++++++++++++++
...test_tir_transform_force_narrow_index_to_i32.py | 14 ++
10 files changed, 366 insertions(+), 106 deletions(-)
diff --git a/include/tvm/s_tir/transform.h b/include/tvm/s_tir/transform.h
index 0a92264c4c..5149a5fe28 100644
--- a/include/tvm/s_tir/transform.h
+++ b/include/tvm/s_tir/transform.h
@@ -375,6 +375,19 @@ TVM_DLL Pass DecorateDeviceScope();
*/
TVM_DLL Pass UseAssumeToReduceBranches();
+/*!
+ * \brief Force to narrow down indexing expressions and integer buffers to
int32 dtype in
+ * functions that may still contain S-TIR blocks.
+ *
+ * Unlike tirx::transform::ForceNarrowIndexToInt32, this pass also rewrites
block iterators,
+ * block access regions, and match buffer regions, so it can run on scheduled
functions before
+ * block lowering.
+ *
+ * \return The pass.
+ * \note This pass should not be used in default cases.
+ */
+TVM_DLL Pass ForceNarrowIndexToInt32();
+
} // namespace transform
} // namespace s_tir
} // namespace tvm
diff --git a/include/tvm/tirx/transform.h b/include/tvm/tirx/transform.h
index edcc05f37b..c37da8a235 100644
--- a/include/tvm/tirx/transform.h
+++ b/include/tvm/tirx/transform.h
@@ -217,6 +217,9 @@ TVM_DLL Pass NarrowDataType(int target_bits);
/*!
* \brief Force to narrow down indexing expressions and integer buffers to
int32 dtype.
*
+ * The function must not contain S-TIR blocks. Use
s_tir::transform::ForceNarrowIndexToInt32
+ * before block lowering.
+ *
* \return The pass.
* \note This pass should not be used in default cases.
*/
diff --git a/python/tvm/s_tir/transform/transform.py
b/python/tvm/s_tir/transform/transform.py
index 6bbaeb0f33..09e04b0a15 100644
--- a/python/tvm/s_tir/transform/transform.py
+++ b/python/tvm/s_tir/transform/transform.py
@@ -512,3 +512,22 @@ def UseAssumeToReduceBranches():
The result pass
"""
return _ffi_api.UseAssumeToReduceBranches() # type: ignore
+
+
+def ForceNarrowIndexToInt32():
+ """Force narrow down indexing expressions and integer buffers to int32
dtype.
+
+ Unlike :py:func:`tvm.tirx.transform.ForceNarrowIndexToInt32`, this pass
also rewrites block
+ iterators, block access regions, and match buffer regions, so it can run
on scheduled
+ functions before block lowering.
+
+ Returns
+ -------
+ fpass : tvm.transform.Pass
+ The result pass
+
+ Note
+ ----
+ This pass should not be used in default cases.
+ """
+ return _ffi_api.ForceNarrowIndexToInt32() # type: ignore
diff --git a/python/tvm/tirx/transform/transform.py
b/python/tvm/tirx/transform/transform.py
index 66dc7f3b5d..bf2e9ba097 100644
--- a/python/tvm/tirx/transform/transform.py
+++ b/python/tvm/tirx/transform/transform.py
@@ -357,6 +357,9 @@ def NarrowDataType(target_bits: int):
def ForceNarrowIndexToInt32():
"""Force narrow down indexing expressions and integer buffers to int32
dtype.
+ The function must not contain S-TIR blocks. Use
+ :py:func:`tvm.s_tir.transform.ForceNarrowIndexToInt32` before block
lowering.
+
Returns
-------
fpass : tvm.transform.Pass
diff --git a/src/s_tir/transform/force_narrow_index_to_i32.cc
b/src/s_tir/transform/force_narrow_index_to_i32.cc
new file mode 100644
index 0000000000..9afbddada3
--- /dev/null
+++ b/src/s_tir/transform/force_narrow_index_to_i32.cc
@@ -0,0 +1,78 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied. See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+/*!
+ * \file force_narrow_index_to_i32.cc
+ * \brief Force narrow down indexing expressions and integer buffers to int32
dtype in functions
+ * that still contain S-TIR blocks.
+ * \note This pass is not used in default cases.
+ */
+
+#include "../../tirx/transform/force_narrow_index_to_i32.h"
+
+#include <tvm/ffi/reflection/registry.h>
+#include <tvm/s_tir/transform.h>
+#include <tvm/tirx/transform.h>
+
+#include "../ir/data_type_rewriter.h"
+
+namespace tvm {
+namespace s_tir {
+using namespace tvm::tirx;
+
+class Int32DTypeNarrower : public
Int32DTypeNarrowerBase<IndexDataTypeNormalizer> {
+ public:
+ using Int32DTypeNarrowerBase::Mutate;
+ using Int32DTypeNarrowerBase::Mutate_;
+ static PrimFunc RewriteDataType(PrimFunc func) {
+ CheckBufferParams(func);
+ auto narrower = ffi::make_object<Int32DTypeNarrower>(func);
+ return narrower->Rewrite(func);
+ }
+
+ explicit Int32DTypeNarrower(PrimFunc func) :
Int32DTypeNarrowerBase(std::move(func)) {}
+
+ private:
+ UnchangedOr<Stmt> Mutate_(const SBlockNode* op, InplaceMode inplace_mode)
final {
+ auto result = IndexDataTypeNormalizer::Mutate_(op, inplace_mode);
+ auto block =
std::move(result).ValueOrUnchanged(ffi::GetRef<Stmt>(op)).as_or_throw<SBlock>();
+ for (const BufferVar& buf : block->alloc_buffers) {
+ CheckAllocatedBuffer(buf);
+ }
+ return block;
+ }
+};
+
+namespace transform {
+
+Pass ForceNarrowIndexToInt32() {
+ auto pass_func = [](PrimFunc f, IRModule m, PassContext ctx) {
+ return Int32DTypeNarrower::RewriteDataType(std::move(f));
+ };
+ return CreatePrimFuncPass(pass_func, 0, "s_tir.ForceNarrowIndexToInt32", {});
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ refl::GlobalDef().def("s_tir.transform.ForceNarrowIndexToInt32",
ForceNarrowIndexToInt32);
+}
+
+} // namespace transform
+} // namespace s_tir
+} // namespace tvm
diff --git a/src/tirx/transform/force_narrow_index_to_i32.cc
b/src/tirx/transform/force_narrow_index_to_i32.cc
index 5692dd7872..00382ab84a 100644
--- a/src/tirx/transform/force_narrow_index_to_i32.cc
+++ b/src/tirx/transform/force_narrow_index_to_i32.cc
@@ -23,71 +23,33 @@
* \note This pass is not used in default cases.
*/
+#include "force_narrow_index_to_i32.h"
+
#include <tvm/ffi/cast.h>
#include <tvm/ffi/reflection/registry.h>
-#include <tvm/tirx/op.h>
+#include <tvm/s_tir/stmt.h>
+#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
-#include "../ir/data_type_rewriter.h"
-
namespace tvm {
namespace tirx {
-using namespace tvm::prim;
-class Int32DTypeNarrower : public IndexDataTypeNormalizer {
+class Int32DTypeNarrower : public
Int32DTypeNarrowerBase<IndexDataTypeNormalizer> {
public:
- using IndexDataTypeNormalizer::Mutate;
- using IndexDataTypeNormalizer::Mutate_;
static PrimFunc RewriteDataType(PrimFunc func) {
- // Check if the integer parameter buffers have dtype other than int32.
- for (const Var& param : func->params) {
- if (auto buffer = param.as<BufferVar>();
- buffer && buffer.value()->dtype.MatchesCode(DLDataTypeCode::kDLInt)
&&
- buffer.value()->dtype.bits() > 32) {
- TVM_FFI_THROW(InternalError) << "The buffer parameter " <<
buffer.value() << " has dtype "
- << buffer.value()->dtype << ". The
function is " << func;
- }
+ // The TIRX normalizer does not rewrite S-TIR block iterators, regions, or
match buffers, so
+ // narrowing a function that still contains blocks would leave their index
types inconsistent.
+ if (ContainsNode<s_tir::SBlockRealizeNode>(func->body)) {
+ TVM_FFI_THROW(ValueError)
+ << "tirx.transform.ForceNarrowIndexToInt32 requires a function
without S-TIR blocks. "
+ << "Use s_tir.transform.ForceNarrowIndexToInt32 before block
lowering.";
}
-
+ CheckBufferParams(func);
auto narrower = ffi::make_object<Int32DTypeNarrower>(func);
return narrower->Rewrite(func);
}
- public:
- explicit Int32DTypeNarrower(PrimFunc func)
- : IndexDataTypeNormalizer(PrimType::Int(32)), func_(std::move(func)) {}
-
- private:
- bool ShouldClampShiftAmounts() const final { return true; }
-
- UnchangedOr<PrimExpr> Mutate_(const IntImmNode* op, InplaceMode
inplace_mode) final {
- // ignore the enabled condition and always rewrite i64
- if (op->ty.as_or_throw<PrimType>() == PrimType::Int(64)) {
- TVM_FFI_ICHECK_LE(op->value,
max_value(target_data_type_).as_or_throw<IntImm>()->value);
- return IntImm::Int32(op->value);
- }
- return ffi::Unchanged();
- }
-
- UnchangedOr<Stmt> Mutate_(const AllocBufferNode* op, InplaceMode
inplace_mode) final {
- auto result = IndexDataTypeNormalizer::Mutate_(op, inplace_mode);
- auto alloc =
-
std::move(result).ValueOrUnchanged(ffi::GetRef<Stmt>(op)).as_or_throw<AllocBuffer>();
- const BufferVar& buf = alloc->buffer;
- // Scalar assignments in TVMScript use local scalar storage. Keep its
explicit
- // dtype (e.g. an int64 opaque call result) and cast at narrowed index
uses.
- // IsScalar checks the scalar layout contract, not merely the allocation
size.
- bool is_local_scalar = buf.scope() == "local" && buf.IsScalar();
- if (!is_local_scalar && buf->dtype.MatchesCode(DLDataTypeCode::kDLInt) &&
- buf->dtype.bits() > 32) {
- TVM_FFI_THROW(InternalError)
- << "The buffer " << buf << " allocated in the function has dtype "
<< buf->dtype
- << ". The function is " << func_;
- }
- return alloc;
- }
-
- PrimFunc func_;
+ explicit Int32DTypeNarrower(PrimFunc func) :
Int32DTypeNarrowerBase(std::move(func)) {}
};
PrimFunc ForceNarrowIndexToInt32(PrimFunc func) {
diff --git a/src/tirx/transform/force_narrow_index_to_i32.cc
b/src/tirx/transform/force_narrow_index_to_i32.h
similarity index 60%
copy from src/tirx/transform/force_narrow_index_to_i32.cc
copy to src/tirx/transform/force_narrow_index_to_i32.h
index 5692dd7872..eba75587d2 100644
--- a/src/tirx/transform/force_narrow_index_to_i32.cc
+++ b/src/tirx/transform/force_narrow_index_to_i32.h
@@ -18,28 +18,39 @@
*/
/*!
- * \file force_narrow_index_to_i32.cc
- * \brief Force narrow down indexing expressions and integer buffers to int32
dtype.
- * \note This pass is not used in default cases.
+ * \file force_narrow_index_to_i32.h
+ * \brief Narrowing rules shared by the TIRX and S-TIR ForceNarrowIndexToInt32
passes.
*/
+#ifndef TVM_TIR_TRANSFORM_FORCE_NARROW_INDEX_TO_I32_H_
+#define TVM_TIR_TRANSFORM_FORCE_NARROW_INDEX_TO_I32_H_
-#include <tvm/ffi/cast.h>
-#include <tvm/ffi/reflection/registry.h>
#include <tvm/tirx/op.h>
-#include <tvm/tirx/transform.h>
+
+#include <utility>
#include "../ir/data_type_rewriter.h"
namespace tvm {
namespace tirx {
-using namespace tvm::prim;
-class Int32DTypeNarrower : public IndexDataTypeNormalizer {
+/*!
+ * \brief Force index expressions and integer buffers to int32.
+ * \tparam Normalizer The index normalizer that determines which statements
are traversed:
+ * tirx::IndexDataTypeNormalizer for lowered TIR,
s_tir::IndexDataTypeNormalizer for
+ * TIR that still contains S-TIR blocks.
+ */
+template <typename Normalizer>
+class Int32DTypeNarrowerBase : public Normalizer {
public:
- using IndexDataTypeNormalizer::Mutate;
- using IndexDataTypeNormalizer::Mutate_;
- static PrimFunc RewriteDataType(PrimFunc func) {
- // Check if the integer parameter buffers have dtype other than int32.
+ using Normalizer::Mutate;
+ using Normalizer::Mutate_;
+
+ protected:
+ explicit Int32DTypeNarrowerBase(PrimFunc func)
+ : Normalizer(PrimType::Int(32)), func_(std::move(func)) {}
+
+ /*! \brief Reject integer buffer parameters wider than int32. */
+ static void CheckBufferParams(const PrimFunc& func) {
for (const Var& param : func->params) {
if (auto buffer = param.as<BufferVar>();
buffer && buffer.value()->dtype.MatchesCode(DLDataTypeCode::kDLInt)
&&
@@ -48,66 +59,48 @@ class Int32DTypeNarrower : public IndexDataTypeNormalizer {
<< buffer.value()->dtype << ". The
function is " << func;
}
}
-
- auto narrower = ffi::make_object<Int32DTypeNarrower>(func);
- return narrower->Rewrite(func);
}
- public:
- explicit Int32DTypeNarrower(PrimFunc func)
- : IndexDataTypeNormalizer(PrimType::Int(32)), func_(std::move(func)) {}
+ /*! \brief Reject allocated integer buffers wider than int32. */
+ void CheckAllocatedBuffer(const BufferVar& buf) const {
+ // Scalar assignments in TVMScript use local scalar storage. Keep its
explicit
+ // dtype (e.g. an int64 opaque call result) and cast at narrowed index
uses.
+ // IsScalar checks the scalar layout contract, not merely the allocation
size.
+ bool is_local_scalar = buf.scope() == "local" && buf.IsScalar();
+ if (!is_local_scalar && buf->dtype.MatchesCode(DLDataTypeCode::kDLInt) &&
+ buf->dtype.bits() > 32) {
+ TVM_FFI_THROW(InternalError)
+ << "The buffer " << buf << " allocated in the function has dtype "
<< buf->dtype
+ << ". The function is " << func_;
+ }
+ }
- private:
bool ShouldClampShiftAmounts() const final { return true; }
UnchangedOr<PrimExpr> Mutate_(const IntImmNode* op, InplaceMode
inplace_mode) final {
// ignore the enabled condition and always rewrite i64
if (op->ty.as_or_throw<PrimType>() == PrimType::Int(64)) {
- TVM_FFI_ICHECK_LE(op->value,
max_value(target_data_type_).as_or_throw<IntImm>()->value);
+ TVM_FFI_ICHECK_LE(
+ op->value,
+ prim::max_value(this->target_data_type_).template
as_or_throw<IntImm>()->value);
return IntImm::Int32(op->value);
}
return ffi::Unchanged();
}
UnchangedOr<Stmt> Mutate_(const AllocBufferNode* op, InplaceMode
inplace_mode) final {
- auto result = IndexDataTypeNormalizer::Mutate_(op, inplace_mode);
- auto alloc =
-
std::move(result).ValueOrUnchanged(ffi::GetRef<Stmt>(op)).as_or_throw<AllocBuffer>();
- const BufferVar& buf = alloc->buffer;
- // Scalar assignments in TVMScript use local scalar storage. Keep its
explicit
- // dtype (e.g. an int64 opaque call result) and cast at narrowed index
uses.
- // IsScalar checks the scalar layout contract, not merely the allocation
size.
- bool is_local_scalar = buf.scope() == "local" && buf.IsScalar();
- if (!is_local_scalar && buf->dtype.MatchesCode(DLDataTypeCode::kDLInt) &&
- buf->dtype.bits() > 32) {
- TVM_FFI_THROW(InternalError)
- << "The buffer " << buf << " allocated in the function has dtype "
<< buf->dtype
- << ". The function is " << func_;
- }
+ auto result = Normalizer::Mutate_(op, inplace_mode);
+ auto alloc = std::move(result)
+ .ValueOrUnchanged(ffi::GetRef<Stmt>(op))
+ .template as_or_throw<AllocBuffer>();
+ CheckAllocatedBuffer(alloc->buffer);
return alloc;
}
PrimFunc func_;
};
-PrimFunc ForceNarrowIndexToInt32(PrimFunc func) {
- return Int32DTypeNarrower::RewriteDataType(func);
-}
-
-namespace transform {
-
-Pass ForceNarrowIndexToInt32() {
- auto pass_func = [](PrimFunc f, IRModule m, PassContext ctx) {
- return ForceNarrowIndexToInt32(f);
- };
- return CreatePrimFuncPass(pass_func, 0, "tirx.NarrowDataType", {});
-}
-
-TVM_FFI_STATIC_INIT_BLOCK() {
- namespace refl = tvm::ffi::reflection;
- refl::GlobalDef().def("tirx.transform.ForceNarrowIndexToInt32",
ForceNarrowIndexToInt32);
-}
-
-} // namespace transform
} // namespace tirx
} // namespace tvm
+
+#endif // TVM_TIR_TRANSFORM_FORCE_NARROW_INDEX_TO_I32_H_
diff --git a/tests/python/relax/test_backend_dispatch_sort_scan.py
b/tests/python/relax/test_backend_dispatch_sort_scan.py
index f616a08b61..02aa7e89ff 100644
--- a/tests/python/relax/test_backend_dispatch_sort_scan.py
+++ b/tests/python/relax/test_backend_dispatch_sort_scan.py
@@ -530,7 +530,7 @@ def test_dispatch_cumsum_gpu(target, index_bits):
with tvm.target.Target(target):
mod = DispatchSortScan(index_bits=index_bits)(Module)
if index_bits == 32:
- mod = tirx.transform.ForceNarrowIndexToInt32()(mod)
+ mod = tvm.s_tir.transform.ForceNarrowIndexToInt32()(mod)
ex = tvm.compile(mod, target)
def run_and_check():
@@ -572,7 +572,7 @@ def test_dispatch_cumsum_index_width(target_kind,
index_bits):
assert_structural_equal(mod["gpu_2d_continuous_cumsum"], expected)
if expected_bits == 32:
# This previously failed on Metal with a 2**35 IntImm.
- tirx.transform.ForceNarrowIndexToInt32()(mod)
+ tvm.s_tir.transform.ForceNarrowIndexToInt32()(mod)
@pytest.mark.parametrize("index_bits", [0, 16, 128])
diff --git
a/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py
b/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py
new file mode 100644
index 0000000000..a651efe62b
--- /dev/null
+++
b/tests/python/s_tir/transform/test_s_tir_transform_force_narrow_index_to_i32.py
@@ -0,0 +1,175 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+import pytest
+
+import tvm
+import tvm.testing
+from tvm.s_tir import dlight as dl
+from tvm.script import tirx as T
+from tvm.testing import env
+
+
+def _narrow(func):
+ mod = tvm.IRModule.from_expr(func)
+ return tvm.s_tir.transform.ForceNarrowIndexToInt32()(mod)["main"]
+
+
+def test_block():
+ @T.prim_func(private=True, s_tir=True)
+ def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")):
+ for i in T.serial(0, T.int64(16)):
+ for j in T.serial(0, T.int64(8)):
+ with T.sblock():
+ vi = T.axis.spatial(T.int64(128), i * T.int64(8) + j)
+ B[vi] = A[vi] + T.float32(1)
+
+ @T.prim_func(private=True, s_tir=True)
+ def expected(A: T.Buffer((128,), "float32"), B: T.Buffer((128,),
"float32")):
+ for i in T.serial(0, T.int32(16)):
+ for j in T.serial(0, T.int32(8)):
+ with T.sblock():
+ vi = T.axis.spatial(T.int32(128), i * T.int32(8) + j)
+ B[vi] = A[vi] + T.float32(1)
+
+ tvm.ir.assert_structural_equal(_narrow(before), expected)
+
+
+def test_block_iters_used_only_in_regions():
+ """Blockized blocks use their iterators only in access and match_buffer
regions."""
+
+ @T.prim_func(private=True, s_tir=True)
+ def before(
+ A: T.Buffer((T.int64(16), T.int64(16)), "float32"),
+ B: T.Buffer((T.int64(16), T.int64(16)), "float32"),
+ ):
+ for i_o, j_o in T.grid(T.int64(2), T.int64(2)):
+ with T.sblock("tile_o"):
+ vi_o, vj_o = T.axis.remap("SS", [i_o, j_o])
+ T.reads(
+ A[
+ vi_o * T.int64(8) : vi_o * T.int64(8) + T.int64(8),
+ vj_o * T.int64(8) : vj_o * T.int64(8) + T.int64(8),
+ ]
+ )
+ T.writes(
+ B[
+ vi_o * T.int64(8) : vi_o * T.int64(8) + T.int64(8),
+ vj_o * T.int64(8) : vj_o * T.int64(8) + T.int64(8),
+ ]
+ )
+ A_tile = T.match_buffer(
+ A[
+ vi_o * T.int64(8) : vi_o * T.int64(8) + T.int64(8),
+ vj_o * T.int64(8) : vj_o * T.int64(8) + T.int64(8),
+ ],
+ (T.int64(8), T.int64(8)),
+ offset_factor=1,
+ )
+ B_tile = T.match_buffer(
+ B[
+ vi_o * T.int64(8) : vi_o * T.int64(8) + T.int64(8),
+ vj_o * T.int64(8) : vj_o * T.int64(8) + T.int64(8),
+ ],
+ (T.int64(8), T.int64(8)),
+ offset_factor=1,
+ )
+ for i_i, j_i in T.grid(T.int64(8), T.int64(8)):
+ with T.sblock("tile"):
+ vi_i, vj_i = T.axis.remap("SS", [i_i, j_i])
+ B_tile[vi_i, vj_i] = A_tile[vi_i, vj_i] + T.float32(1)
+
+ @T.prim_func(private=True, s_tir=True)
+ def expected(A: T.Buffer((16, 16), "float32"), B: T.Buffer((16, 16),
"float32")):
+ for i_o, j_o in T.grid(2, 2):
+ with T.sblock("tile_o"):
+ vi_o, vj_o = T.axis.remap("SS", [i_o, j_o])
+ T.reads(A[vi_o * 8 : vi_o * 8 + 8, vj_o * 8 : vj_o * 8 + 8])
+ T.writes(B[vi_o * 8 : vi_o * 8 + 8, vj_o * 8 : vj_o * 8 + 8])
+ A_tile = T.match_buffer(
+ A[vi_o * 8 : vi_o * 8 + 8, vj_o * 8 : vj_o * 8 + 8], (8,
8), offset_factor=1
+ )
+ B_tile = T.match_buffer(
+ B[vi_o * 8 : vi_o * 8 + 8, vj_o * 8 : vj_o * 8 + 8], (8,
8), offset_factor=1
+ )
+ for i_i, j_i in T.grid(8, 8):
+ with T.sblock("tile"):
+ vi_i, vj_i = T.axis.remap("SS", [i_i, j_i])
+ B_tile[vi_i, vj_i] = A_tile[vi_i, vj_i] + T.float32(1)
+
+ tvm.ir.assert_structural_equal(_narrow(before), expected)
+
+
+def test_fail_on_buffer_param():
+ @T.prim_func(private=True, s_tir=True)
+ def func(A: T.Buffer((128,), "int64"), B: T.Buffer((128,), "int64")):
+ for i in T.serial(0, 16):
+ for j in T.serial(0, 8):
+ with T.sblock():
+ vi = T.axis.spatial(128, i * 8 + j)
+ B[vi] = A[vi] + T.int64(1)
+
+ with pytest.raises(RuntimeError):
+ _narrow(func)
+
+
+def test_fail_on_block_alloc_buffer():
+ @T.prim_func(private=True, s_tir=True)
+ def func(A: T.Buffer((128,), "int32"), B: T.Buffer((128,), "int32")):
+ C = T.sblock_alloc_buffer((128,), "int64")
+ for i in T.serial(0, 16):
+ for j in T.serial(0, 8):
+ with T.sblock():
+ vi = T.axis.spatial(128, i * 8 + j)
+ C[vi] = T.cast(A[vi], "int64") + T.int64(1)
+ for i in T.serial(0, 16):
+ for j in T.serial(0, 8):
+ with T.sblock():
+ vi = T.axis.spatial(128, i * 8 + j)
+ B[vi] = T.cast(C[vi] + T.int64(1), "int32")
+
+ with pytest.raises(RuntimeError):
+ _narrow(func)
+
+
[email protected](not env.has_llvm(), reason="need llvm for the host module")
+def test_metal_simdgroup_matmul_builds():
+ """Narrowing a DLight-scheduled Metal matmul keeps its tensorized blocks
consistent."""
+
+ @T.prim_func(s_tir=True)
+ def main(
+ var_A: T.handle, B: T.Buffer((T.int64(256), T.int64(256)), "float16"),
var_C: T.handle
+ ):
+ n = T.int64()
+ A = T.match_buffer(var_A, (T.int64(1), n, T.int64(256)), "float16")
+ C = T.match_buffer(var_C, (T.int64(1), n, T.int64(256)), "float16")
+ for i0, i1, i2, k in T.grid(T.int64(1), n, T.int64(256), T.int64(256)):
+ with T.sblock("NT_matmul"):
+ v0, v1, v2, vk = T.axis.remap("SSSR", [i0, i1, i2, k])
+ with T.init():
+ C[v0, v1, v2] = T.float16(0)
+ C[v0, v1, v2] = C[v0, v1, v2] + A[v0, v1, vk] * B[v2, vk]
+
+ target = tvm.target.Target("metal", host="llvm")
+ with target:
+ mod = dl.ApplyDefaultSchedule(dl.gpu.Matmul())(tvm.IRModule({"main":
main}))
+ assert "metal.simdgroup" in mod.script()
+ mod = tvm.s_tir.transform.ForceNarrowIndexToInt32()(mod)
+ tvm.tirx.build(mod, target=target)
+
+
+if __name__ == "__main__":
+ tvm.testing.main()
diff --git
a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py
b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py
index bb99a48363..f4a0d68e8a 100644
---
a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py
+++
b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py
@@ -202,6 +202,20 @@ def test_block():
tvm.ir.assert_structural_equal(func, _lower_blocks(expected))
+def test_reject_blocks():
+ @T.prim_func(private=True, s_tir=True)
+ def before(A: T.Buffer((128,), "float32"), B: T.Buffer((128,), "float32")):
+ for i in T.serial(0, T.int64(16)):
+ for j in T.serial(0, T.int64(8)):
+ with T.sblock():
+ vi = T.axis.spatial(T.int64(128), i * T.int64(8) + j)
+ B[vi] = A[vi] + T.float32(1)
+
+ mod = tvm.IRModule.from_expr(before)
+ with pytest.raises(ValueError, match="requires a function without S-TIR
blocks"):
+ tvm.tirx.transform.ForceNarrowIndexToInt32()(mod)
+
+
def test_i16_buffer():
@T.prim_func(private=True, s_tir=True)
def before(A: T.Buffer((128,), "int16"), B: T.Buffer((128,), "int16")):