Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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 @@ -167,5 +167,53 @@ def kernel(A: T.Tensor((2,), "int32")):
torch.testing.assert_close(result.cpu(), torch.tensor([2, 4], dtype=torch.int32), rtol=0, atol=0)


@pytest.mark.parametrize("constant", [10, T.int32(10)])
def test_constant_rebind_inside_a_region_is_rejected(constant):
"""The expression form of this rebind is rejected already.

A constant right-hand side took the fast path in `bind`, which dropped the
binding record instead: the branch was lost, the constant replaced the value at
trace time, and every later read saw it unconditionally.
"""
with pytest.raises(RuntimeError, match="outside its defining region"):

@T.prim_func
def kernel(A: T.Tensor((4,), "int32"), Out: T.Tensor((4,), "int32")):
with T.Kernel(1, threads=1):
for i in T.serial(4):
val = A[i]
if val < 10:
val = constant
Out[i] = val


def test_constant_rebind_at_the_same_level_is_still_accepted():
"""Control: nothing is conditional here, so there is nothing to reject."""

@T.prim_func
def kernel(Out: T.Tensor((2,), "int32")):
with T.Kernel(1, threads=1):
value = 1
value = 2
Out[0] = value

tilelang.compile(kernel,>


def test_constant_stays_reusable_once_its_region_closed():
"""Control: a constant is not a TIR binding, so it does not expire with its region."""

@T.prim_func
def kernel(Out: T.Tensor((2,), "int32")):
with T.Kernel(1, threads=1):
for i in T.serial(1):
reused = 7
Out[i] = reused
for j in T.serial(1, 2):
Out[j] = reused

tilelang.compile(kernel,>


if __name__ == "__main__":
tilelang.testing.main()
29 changes: 29 additions & 0 deletions tilelang/language/eager/builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -684,9 +684,11 @@ def bind(self, name, value, annot=BaseBuilder.empty, *, loop_target=False):

# 2. Quick return for trivil types
if isinstance(value, (tuple, list, tvm.ffi.Array, int, float, str)):
self.reject_conditional_constant_rebind(name)
self.name_inside_frame.pop(name, None)
return value
if isinstance(value, tirx.IntImm) and value.dtype == "int32":
self.reject_conditional_constant_rebind(name)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
return value.value
if isinstance(value, (Var, Buffer)):
# Bind TVM Var/Buffer names and also record scope so reusing the same
Expand All @@ -713,6 +715,33 @@ def bind(self, name, value, annot=BaseBuilder.empty, *, loop_target=False):
self.name_inside_frame[name] = self.frames[frame]
return res

def reject_conditional_constant_rebind(self, name: str | None) -> None:
"""Reject a constant assigned to a name bound in an enclosing region.

The constant fast paths in `bind` take an `int`/`float`/`str`/`IntImm` and
return it directly. Assigning one to a name that already has a live
binding in an enclosing region was accepted silently: the constant
replaced the value at trace time, the enclosing `if` was dropped, and
every later read saw the constant unconditionally. The expression form of
the same rebind is already rejected, so a clamp written in the natural way
returned a wrong tensor with no diagnostic.

A name whose binding is in a region that has already closed is left alone,
so a constant stays reusable once the region that introduced it is gone.
"""
if name is None or name == "_":
return
bound_in = self.name_inside_frame.get(name)
if bound_in is None or bound_in not in self.frames:
return
innermost = self.find_frame_idx(TIR_VAR_SCOPE_FRAME)
if innermost is None or self.frames[innermost] is bound_in:
return
raise RuntimeError(
f"Immutable variable `{name}` is used outside its defining region!\n"
f"variable `{name}` is defined in frame: {self.name_inside_frame[name]}, current frames: {self.frames}."
)

def binding_expired(self, name: str | None) -> bool:
"""Whether `name` was last bound inside a TIR region that has since closed."""
frame = self.name_inside_frame.get(name)
Expand Down
Loading