Skip to content

Commit 72d6348

Browse files
committed
[BugFix] Reject a constant rebound inside a nested region
`Builder.bind` has a fast path for a constant right-hand side (an int, float, str or int32 IntImm) that returns the value directly and drops the name's binding record. Assigning a constant to a name that already has a live binding in an enclosing region therefore went through unremarked: the constant replaced the value at trace time, the enclosing `if` was dropped, and every later read saw the constant unconditionally. for i in T.serial(4): val = A[i] if val < 10: val = 10 if val > 100: val = 100 Out[i] = val With A = [5, 50, 150, 10] this returned [10, 10, 10, 10] instead of [10, 50, 100, 10]. The expression form of the same rebind is already rejected with "Immutable variable `val` is used outside its defining region!", so the same illegal program was accepted or rejected depending only on the form of the right-hand side. The fast path now rejects that rebind with the same message. A constant whose region has already closed stays reusable, so nothing that a closed scope introduced starts expiring.
1 parent 994b44e commit 72d6348

2 files changed

Lines changed: 77 additions & 0 deletions

File tree

‎testing/python/language/test_tilelang_language_loop_target_binding.py‎

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -167,5 +167,53 @@ def kernel(A: T.Tensor((2,), "int32")):
167167
torch.testing.assert_close(result.cpu(), torch.tensor([2, 4], dtype=torch.int32), rtol=0, atol=0)
168168

169169

170+
@pytest.mark.parametrize("constant", [10, T.int32(10)])
171+
def test_constant_rebind_inside_a_region_is_rejected(constant):
172+
"""The expression form of this rebind is rejected already.
173+
174+
A constant right-hand side took the fast path in `bind`, which dropped the
175+
binding record instead: the branch was lost, the constant replaced the value at
176+
trace time, and every later read saw it unconditionally.
177+
"""
178+
with pytest.raises(RuntimeError, match="outside its defining region"):
179+
180+
@T.prim_func
181+
def kernel(A: T.Tensor((4,), "int32"), Out: T.Tensor((4,), "int32")):
182+
with T.Kernel(1, threads=1):
183+
for i in T.serial(4):
184+
val = A[i]
185+
if val < 10:
186+
val = constant
187+
Out[i] = val
188+
189+
190+
def test_constant_rebind_at_the_same_level_is_still_accepted():
191+
"""Control: nothing is conditional here, so there is nothing to reject."""
192+
193+
@T.prim_func
194+
def kernel(Out: T.Tensor((2,), "int32")):
195+
with T.Kernel(1, threads=1):
196+
value = 1
197+
value = 2
198+
Out[0] = value
199+
200+
tilelang.compile(kernel, target="cuda")
201+
202+
203+
def test_constant_stays_reusable_once_its_region_closed():
204+
"""Control: a constant is not a TIR binding, so it does not expire with its region."""
205+
206+
@T.prim_func
207+
def kernel(Out: T.Tensor((2,), "int32")):
208+
with T.Kernel(1, threads=1):
209+
for i in T.serial(1):
210+
reused = 7
211+
Out[i] = reused
212+
for j in T.serial(1, 2):
213+
Out[j] = reused
214+
215+
tilelang.compile(kernel, target="cuda")
216+
217+
170218
if __name__ == "__main__":
171219
tilelang.testing.main()

‎tilelang/language/eager/builder.py‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -684,9 +684,11 @@ def bind(self, name, value, annot=BaseBuilder.empty, *, loop_target=False):
684684

685685
# 2. Quick return for trivil types
686686
if isinstance(value, (tuple, list, tvm.ffi.Array, int, float, str)):
687+
self.reject_conditional_constant_rebind(name)
687688
self.name_inside_frame.pop(name, None)
688689
return value
689690
if isinstance(value, tirx.IntImm) and value.dtype == "int32":
691+
self.reject_conditional_constant_rebind(name)
690692
return value.value
691693
if isinstance(value, (Var, Buffer)):
692694
# Bind TVM Var/Buffer names and also record scope so reusing the same
@@ -713,6 +715,33 @@ def bind(self, name, value, annot=BaseBuilder.empty, *, loop_target=False):
713715
self.name_inside_frame[name] = self.frames[frame]
714716
return res
715717

718+
def reject_conditional_constant_rebind(self, name: str | None) -> None:
719+
"""Reject a constant assigned to a name bound in an enclosing region.
720+
721+
The constant fast paths in `bind` take an `int`/`float`/`str`/`IntImm` and
722+
return it directly. Assigning one to a name that already has a live
723+
binding in an enclosing region was accepted silently: the constant
724+
replaced the value at trace time, the enclosing `if` was dropped, and
725+
every later read saw the constant unconditionally. The expression form of
726+
the same rebind is already rejected, so a clamp written in the natural way
727+
returned a wrong tensor with no diagnostic.
728+
729+
A name whose binding is in a region that has already closed is left alone,
730+
so a constant stays reusable once the region that introduced it is gone.
731+
"""
732+
if name is None or name == "_":
733+
return
734+
bound_in = self.name_inside_frame.get(name)
735+
if bound_in is None or bound_in not in self.frames:
736+
return
737+
innermost = self.find_frame_idx(TIR_VAR_SCOPE_FRAME)
738+
if innermost is None or self.frames[innermost] is bound_in:
739+
return
740+
raise RuntimeError(
741+
f"Immutable variable `{name}` is used outside its defining region!\n"
742+
f"variable `{name}` is defined in frame: {self.name_inside_frame[name]}, current frames: {self.frames}."
743+
)
744+
716745
def binding_expired(self, name: str | None) -> bool:
717746
"""Whether `name` was last bound inside a TIR region that has since closed."""
718747
frame = self.name_inside_frame.get(name)

0 commit comments

Comments
 (0)