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

Reply via email to