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
Original file line number Diff line number Diff line change
Expand Up @@ -66,3 +66,41 @@ def main(
assert f"Requested coalesced_width={coalesced_width}" in warnings
assert "using 4 instead" in warnings
assert "float4" in artifact.kernel_source


@pytest.mark.parametrize("coalesced_width", [2, 4])
def test_parallel_coalesced_width_integer(coalesced_width):
# Regression test for issue: T.Parallel(coalesced_width=<int>) should accept
# raw Python integers without crashing LayoutInference with "coalesced_width should be an IntImmNode".
"cuda", "arch": "sm_80"})
m, n = 128, 128

@T.prim_func
def main(
A: T.Tensor((m, n), T.float32),
B: T.Tensor((m, n), T.float32),
):
with T.Kernel(1, threads=128):
for i, j in T.Parallel(m, n, coalesced_width=coalesced_width):
B[i, j] = A[i, j]

artifact = _lower_without_device_compile(main, target)
assert artifact.kernel_source


def test_parallel_coalesced_width_annotation_dict():
# Regression test: T.Parallel with annotations={"coalesced_width": <int>}
"cuda", "arch": "sm_80"})
m, n = 128, 128

@T.prim_func
def main(
A: T.Tensor((m, n), T.float32),
B: T.Tensor((m, n), T.float32),
):
with T.Kernel(1, threads=128):
for i, j in T.Parallel(m, n, annotations={"coalesced_width": 4}):
B[i, j] = A[i, j]

artifact = _lower_without_device_compile(main, target)
assert artifact.kernel_source
4 changes: 4 additions & 0 deletions tilelang/language/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,10 @@ def Parallel(
merged_annotations: dict[str, Any] = dict(annotations) if annotations is not None else {}
if coalesced_width is not None:
merged_annotations["coalesced_width"] = coalesced_width
if "coalesced_width" in merged_annotations:
cw = merged_annotations["coalesced_width"]
if not isinstance(cw, tirx.PrimExpr):
merged_annotations["coalesced_width"] = IntImm("int32", int(cw))
if loop_layout is not None:
# Pass through to C++ as the standard parallel loop layout key.
# The builder will attach it only on the outermost parallel loop.
Expand Down