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: