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

tlopex 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 5b4b75357a [DLight][GPU] Skip fallback scheduling for zero extents 
(#20394)
5b4b75357a is described below

commit 5b4b75357a108f83e7b5dfa7787ca1b5eff762ad
Author: Nanmur <[email protected]>
AuthorDate: Tue Sep 22 11:13:09 2026 +0800

    [DLight][GPU] Skip fallback scheduling for zero extents (#20394)
    
    ## Summary
    - Skip the GPU fallback schedule rule when a PrimFunc contains a
    statically zero-extent loop.
    - Add a regression test for an elementwise `float32[4, 0]` function.
    
    This prevents Fallback from reaching inline/fuse transformations that
    construct invalid zero-extent index arithmetic.
    
    ## Testing
    - `python -m pytest
    
tests/python/s_tir/dlight/test_gpu_fallback.py::test_fallback_skips_zero_extent_spatial
    -q`
    - `ruff check python/tvm/s_tir/dlight/gpu/fallback.py
    tests/python/s_tir/dlight/test_gpu_fallback.py`
    
    Fixes #20279
---
 python/tvm/s_tir/dlight/gpu/fallback.py        | 15 +++++++++++++++
 tests/python/s_tir/dlight/test_gpu_fallback.py | 17 +++++++++++++++++
 2 files changed, 32 insertions(+)

diff --git a/python/tvm/s_tir/dlight/gpu/fallback.py 
b/python/tvm/s_tir/dlight/gpu/fallback.py
index e64f8544ec..ebb6998e18 100644
--- a/python/tvm/s_tir/dlight/gpu/fallback.py
+++ b/python/tvm/s_tir/dlight/gpu/fallback.py
@@ -53,6 +53,19 @@ def _has_internal_thread_env(stmt: tirx.Stmt) -> bool:
     return found
 
 
+def _has_zero_extent_loop(stmt: tirx.Stmt) -> bool:
+    """Check whether a statement contains a statically empty loop."""
+    found = False
+
+    def visit_for(node: tirx.For):
+        nonlocal found
+        if isinstance(node.extent, tirx.IntImm) and node.extent.value == 0:
+            found = True
+
+    tvm_ffi.structural_walk(stmt, (tirx.For, visit_for), order="post")
+    return found
+
+
 class Fallback(GPUScheduleRule):
     """
     A fallback schedule rule for all GPU operators. It will try to inline all 
the blocks first,
@@ -67,6 +80,8 @@ class Fallback(GPUScheduleRule):
     ) -> s_tir.Schedule:
         if not isinstance(func, tirx.PrimFunc) or not 
self.is_target_available(target):
             return None
+        if _has_zero_extent_loop(func.body):
+            return None
         max_threads_per_block = base.max_threads_per_block(target)
 
         sch = s_tir.Schedule(func)
diff --git a/tests/python/s_tir/dlight/test_gpu_fallback.py 
b/tests/python/s_tir/dlight/test_gpu_fallback.py
index 76bb99c8a0..74a6de92e2 100644
--- a/tests/python/s_tir/dlight/test_gpu_fallback.py
+++ b/tests/python/s_tir/dlight/test_gpu_fallback.py
@@ -69,6 +69,23 @@ def test_fallback():
     assert_structural_equal(mod, After)
 
 
+def test_fallback_skips_zero_extent_spatial():
+    @I.ir_module(s_tir=True)
+    class Module:
+        @T.prim_func(s_tir=True)
+        def main(A: T.Buffer((4, 0), "float32"), B: T.Buffer((4, 0), 
"float32")):
+            for i, j in T.grid(4, 0):
+                with T.sblock("copy"):
+                    vi, vj = T.axis.remap("SS", [i, j])
+                    B[vi, vj] = A[vi, vj]
+
+    with Target("nvidia/geforce-rtx-3090-ti"):
+        mod = dl.ApplyDefaultSchedule(  # pylint: disable=not-callable
+            dl.gpu.Fallback(),
+        )(Module)
+    assert_structural_equal(mod, Module)
+
+
 def test_fallback_reduction():
     @I.ir_module(s_tir=True)
     class Module:

Reply via email to