Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 21 additions & 5 deletions python/tvm/s_tir/dlight/gpu/reduction.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,10 +107,16 @@ def apply( # pylint: disable=too-many-locals,too-many-branches,too-many-return-
self._sch_inner_reduction(
sch, target, block, c_factor, epilogue, loop_order, s_split_index
)
else:
elif (
self._sch_inner_spatial(
sch, target, block, block_info, c_factor, epilogue, loop_order, s_split_index
)
is None
):
# `_sch_inner_spatial` bails (returns None) on nest shapes it does not model;
# `sch` may already carry partial, unbound edits at that point, so it must not
# be handed back as if scheduling succeeded. Let `Fallback` take over instead.
return None
return sch

def _normalize( # pylint: disable=too-many-branches
Expand Down Expand Up @@ -249,6 +255,7 @@ def _sch_inner_spatial(
s, r, _ = sch.get_loops(block)
len_tx, len_ty = 16, 16
s_factor = [i.dom for i in block_info.iters if i.kind == "S"][-1]

# get perfect spatial factor, spatial factor should be divide the innermost spatial loop so
# that the block after r_factor and be reversed compute at the original scope
while len_tx > 1:
Expand All @@ -268,10 +275,18 @@ def _sch_inner_spatial(
sch.decompose_reduction(rf, r)
# Schedule the write back block
sch.reverse_compute_at(block, bx, preserve_unit_loops=True)
_, r, *s = sch.get_loops(block)
_, r, *s_loops = sch.get_loops(block)

if len([i for i in block_info.iters if i.kind == "S"]) > 2:
return None

if unroll_spatial_factor:
assert len(s) == len(loop_order)
new_order_s = [s[loop_order[i]] for i in range(len(s))]
if len(s_loops) != len(loop_order):
# `loop_order` is indexed by the original block's spatial loops; when
# `reverse_compute_at` regenerates a different number of them the
# premutation is meaningless. Bail rather than mis-permute.
return None
new_order_s = [s_loops[loop_order[i]] for i in range(len(s_loops))]
sch.reorder(*new_order_s)
new_order_s[s_split_index], c = sch.split(
new_order_s[s_split_index], factors=[None, unroll_spatial_factor]
Expand All @@ -280,7 +295,7 @@ def _sch_inner_spatial(
s = sch.fuse(*new_order_s)
sch.reorder(s, c, r)
else:
s = sch.fuse(*s)
s = sch.fuse(*s_loops)
sch.reorder(s, r)
sch.bind(s, "threadIdx.x")
sch.bind(r, "threadIdx.y")
Expand All @@ -303,3 +318,4 @@ def _sch_inner_spatial(
tx, _ = sch.split(sch.fuse(*s), factors=[len_tx, None])
sch.bind(tx, "threadIdx.x")
# pylint: enable=invalid-name
return sch
167 changes: 167 additions & 0 deletions tests/python/s_tir/dlight/test_gpu_reduction.py
Original file line number Diff line number Diff line change
Expand Up @@ -1179,5 +1179,172 @@ def matmul(lv43: T.Buffer((T.int64(1), T.int64(32), T.int64(1)), "float16"), lv4
assert_structural_equal(mod, Before)


def test_reduction_4d_non_contiguous_reduce_axis0():
# Test reduction with kind-based loop classification in _sch_inner_spatial.
# With 3+ spatial dims, `blockIdx.x` alone fully absorbs 2+ of them, so the
# write-back block ends up with 2+ degenerate (extent-1) spatial loops; fusing
# and binding those would require a non-canonical floordiv/subtraction binding.
# `_sch_inner_spatial` detects this and bails, so `Fallback` schedules it instead.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 16, 32), "float32"), A_red: T.Buffer((8, 16, 32), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1, ax2, k in T.grid(8, 16, 32, 4):
with T.sblock("A_red"):
v_ax0, v_ax1, v_ax2, v_k = T.axis.remap("SSSR", [ax0, ax1, ax2, k])
T.reads(A[v_k, v_ax0, v_ax1, v_ax2])
T.writes(A_red[v_ax0, v_ax1, v_ax2])
with T.init():
A_red[v_ax0, v_ax1, v_ax2] = T.float32(0)
A_red[v_ax0, v_ax1, v_ax2] = A_red[v_ax0, v_ax1, v_ax2] + A[v_k, v_ax0, v_ax1, v_ax2]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule( # pylint: disable=not-callable
dl.gpu.Reduction(), dl.gpu.Fallback()
)(Before)
assert mod["main"].attrs["tirx.is_scheduled"] == 1


def test_reduction_4d_non_contiguous_reduce_axis1():
# Test reduction with kind-based loop classification on 4D tensor, reduce axis 1.
# Same 3+-spatial-dims shape as axis0 above: `_sch_inner_spatial` bails on the
# 2+ degenerate write-back loops and `Fallback` schedules it instead.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 16, 32), "float32"), A_red: T.Buffer((4, 16, 32), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1, ax2, k in T.grid(4, 16, 32, 8):
with T.sblock("A_red"):
v_ax0, v_ax1, v_ax2, v_k = T.axis.remap("SSSR", [ax0, ax1, ax2, k])
T.reads(A[v_ax0, v_k, v_ax1, v_ax2])
T.writes(A_red[v_ax0, v_ax1, v_ax2])
with T.init():
A_red[v_ax0, v_ax1, v_ax2] = T.float32(0)
A_red[v_ax0, v_ax1, v_ax2] = A_red[v_ax0, v_ax1, v_ax2] + A[v_ax0, v_k, v_ax1, v_ax2]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule( # pylint: disable=not-callable
dl.gpu.Reduction(), dl.gpu.Fallback()
)(Before)
assert mod["main"].attrs["tirx.is_scheduled"] == 1


def test_reduction_4d_non_contiguous_reduce_axis1_keepdims():
# Test reduction with keepdims=True on non-contiguous axis.
# Unit dimension inserted at reduce position is collapsed by NormalizePrimFunc.
# leaving the same 3+-spatial-dims shape as axis0/axis1 above: `_sch_inner_spatial`
# bails on the 2+ degenerate write-back loops and `Fallback` schedules it instead.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 16, 32), "float32"), A_red: T.Buffer((4, 1, 16, 32), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1u, ax1, ax2, k in T.grid(4, 1, 16, 32, 8):
with T.sblock("A_red"):
v_ax0, v_ax1u, v_ax1, v_ax2, v_k = T.axis.remap("SSSSR", [ax0, ax1u, ax1, ax2, k])
T.reads(A[v_ax0, v_k, v_ax1, v_ax2])
T.writes(A_red[v_ax0, v_ax1u, v_ax1, v_ax2])
with T.init():
A_red[v_ax0, v_ax1u, v_ax1, v_ax2] = T.float32(0)
A_red[v_ax0, v_ax1u, v_ax1, v_ax2] = A_red[v_ax0, v_ax1u, v_ax1, v_ax2] + A[v_ax0, v_k, v_ax1, v_ax2]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule( # pylint: disable=not-callable
dl.gpu.Reduction(), dl.gpu.Fallback()
)(Before)
assert mod["main"].attrs["tirx.is_scheduled"] == 1


def test_reduction_4d_non_contiguous_reduce_axis2():
# Test reduction with kind-based loop classification on 4D tensor, reduce axis 2.
# Same 3+-spatial-dims shape as the other 4D tests above: `_sch_inner_spatial` bails
# on the 2+ degenerate write-back loops and `Fallback` schedules it instead.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 16, 32), "float32"), A_red: T.Buffer((4, 8, 32), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1, ax2, k in T.grid(4, 8, 32, 16):
with T.sblock("A_red"):
v_ax0, v_ax1, v_ax2, v_k = T.axis.remap("SSSR", [ax0, ax1, ax2, k])
T.reads(A[v_ax0, v_ax1, v_k, v_ax2])
T.writes(A_red[v_ax0, v_ax1, v_ax2])
with T.init():
A_red[v_ax0, v_ax1, v_ax2] = T.float32(0)
A_red[v_ax0, v_ax1, v_ax2] = A_red[v_ax0, v_ax1, v_ax2] + A[v_ax0, v_ax1, v_k, v_ax2]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule( # pylint: disable=not-callable
dl.gpu.Reduction(), dl.gpu.Fallback()
)(Before)
assert mod["main"].attrs["tirx.is_scheduled"] == 1


def test_reduction_4d_contiguous_reduce_axis3_regression():
# Regression guard: contiguous (last-axis) reduction on 4D tensor.
# Takes _sch_inner_reduction path (not affected by _sch_inner_spatial changes);
# must keep passing to ensure no unintended side effects.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 16, 32), "float32"), A_red: T.Buffer((4, 8, 16), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1, ax2, k in T.grid(4, 8, 16, 32):
with T.sblock("A_red"):
v_ax0, v_ax1, v_ax2, v_k = T.axis.remap("SSSR", [ax0, ax1, ax2, k])
T.reads(A[v_ax0, v_ax1, v_ax2, v_k])
T.writes(A_red[v_ax0, v_ax1, v_ax2])
with T.init():
A_red[v_ax0, v_ax1, v_ax2] = T.float32(0)
A_red[v_ax0, v_ax1, v_ax2] = A_red[v_ax0, v_ax1, v_ax2] + A[v_ax0, v_ax1, v_ax2, v_k]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule(dl.gpu.Reduction())(Before) # pylint: disable=not-callable
assert mod["main"].attrs["tirx.is_scheduled"] == 1


def test_reduction_3d_non_contiguous_reduce_axis0_regression():
# Regression guard: non-contiguous reduction on 3D tensor with 2 spatial dimensions.
# Takes _sch_inner_spatial path as 4D tests but with fewer spatial dimensions.
# ensures existing working cases remain unaffected by kind-based classification.
# fmt: off
@I.ir_module(s_tir=True)
class Before:
@T.prim_func(s_tir=True)
def main(A: T.Buffer((4, 8, 32), "float32"), A_red: T.Buffer((8, 32), "float32")):
T.func_attr({"tirx.noalias": True})
for ax0, ax1, k in T.grid(8, 32, 4):
with T.sblock("A_red"):
v_ax0, v_ax1, v_k = T.axis.remap("SSR", [ax0, ax1, k])
T.reads(A[v_k, v_ax0, v_ax1])
T.writes(A_red[v_ax0, v_ax1])
with T.init():
A_red[v_ax0, v_ax1] = T.float32(0)
A_red[v_ax0, v_ax1] = A_red[v_ax0, v_ax1] + A[v_k, v_ax0, v_ax1]
# fmt: on

target = Target("nvidia/geforce-rtx-3090-ti")
with target:
mod = dl.ApplyDefaultSchedule(dl.gpu.Reduction())(Before) # pylint: disable=not-callable
assert mod["main"].attrs["tirx.is_scheduled"] == 1


if __name__ == "__main__":
tvm.testing.main()
Loading