Skip to content

Commit 03277bc

Browse files
committed
fix(language): coerce coalesced_width to IntImm in T.Parallel
Coerce Python integers passed to T.Parallel via coalesced_width or annotations['coalesced_width'] into IntImm. Previously, passing a raw integer caused LayoutInference (ParallelOpNode::ComputePlanCandidate) to abort with 'coalesced_width should be an IntImmNode' because TVM FFI does not automatically wrap Python integers into IntImmNode in loop annotation maps. Closes #3013
1 parent 994b44e commit 03277bc

2 files changed

Lines changed: 43 additions & 0 deletions

File tree

‎testing/python/transform/test_tilelang_transform_coalesced_width.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -66,3 +66,42 @@ def main(
6666
assert f"Requested coalesced_width={coalesced_width}" in warnings
6767
assert "using 4 instead" in warnings
6868
assert "float4" in artifact.kernel_source
69+
70+
71+
@pytest.mark.parametrize("coalesced_width", [2, 4])
72+
def test_parallel_coalesced_width_integer(coalesced_width):
73+
# Regression test for issue: T.Parallel(coalesced_width=<int>) should accept
74+
# raw Python integers without crashing LayoutInference with "coalesced_width should be an IntImmNode".
75+
target = tvm.target.Target({"kind": "cuda", "arch": "sm_80"})
76+
m, n = 128, 128
77+
78+
@T.prim_func
79+
def main(
80+
A: T.Tensor((m, n), T.float32),
81+
B: T.Tensor((m, n), T.float32),
82+
):
83+
with T.Kernel(1, threads=128):
84+
for i, j in T.Parallel(m, n, coalesced_width=coalesced_width):
85+
B[i, j] = A[i, j]
86+
87+
artifact = _lower_without_device_compile(main, target)
88+
assert artifact.kernel_source
89+
90+
91+
def test_parallel_coalesced_width_annotation_dict():
92+
# Regression test: T.Parallel with annotations={"coalesced_width": <int>}
93+
target = tvm.target.Target({"kind": "cuda", "arch": "sm_80"})
94+
m, n = 128, 128
95+
96+
@T.prim_func
97+
def main(
98+
A: T.Tensor((m, n), T.float32),
99+
B: T.Tensor((m, n), T.float32),
100+
):
101+
with T.Kernel(1, threads=128):
102+
for i, j in T.Parallel(m, n, annotations={"coalesced_width": 4}):
103+
B[i, j] = A[i, j]
104+
105+
artifact = _lower_without_device_compile(main, target)
106+
assert artifact.kernel_source
107+

‎tilelang/language/loop.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,10 @@ def Parallel(
7272
merged_annotations: dict[str, Any] = dict(annotations) if annotations is not None else {}
7373
if coalesced_width is not None:
7474
merged_annotations["coalesced_width"] = coalesced_width
75+
if "coalesced_width" in merged_annotations:
76+
cw = merged_annotations["coalesced_width"]
77+
if not isinstance(cw, tirx.PrimExpr):
78+
merged_annotations["coalesced_width"] = IntImm("int32", int(cw))
7579
if loop_layout is not None:
7680
# Pass through to C++ as the standard parallel loop layout key.
7781
# The builder will attach it only on the outermost parallel loop.

0 commit comments

Comments
 (0)