Skip to content
Closed
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
@@ -0,0 +1,40 @@
import pytest

import tilelang
import tilelang.language as T


def test_if_rejects_a_non_boolean_buffer_condition_at_the_source():
with pytest.raises(Exception, match="If condition must be a boolean expression, but got int32"):

@T.prim_func
def main(A: T.Tensor((8,), "int32"), B: T.Tensor((8,), "int32")):
with T.Kernel(1, threads=8):
i = T.get_thread_binding()
if A[i]:
B[i] = 1
else:
B[i] = 0


def test_assert_rejects_a_non_boolean_buffer_condition_at_the_source():
with pytest.raises(Exception, match="Assert condition must be a boolean expression, but got int32"):

@T.prim_func
def main(A: T.Tensor((8,), "int32")):
with T.Kernel(1, threads=8):
i = T.get_thread_binding()
assert A[i], "A[i] must be nonzero"


def test_boolean_buffer_conditions_still_lower():
@T.prim_func
def main(A: T.Tensor((8,), "bool"), B: T.Tensor((8,), "int32")):
with T.Kernel(1, threads=8):
i = T.get_thread_binding()
if A[i]:
B[i] = 1
else:
B[i] = 0

tilelang.lower(main,>
6 changes: 5 additions & 1 deletion tilelang/language/parser/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -486,7 +486,9 @@ def visit_if(self: Parser, node: doc.If) -> None:
with self.var_table.with_frame():
predicate = self.eval_expr(node.test)
if isinstance(predicate, (PrimExpr, tvm.tirx.expr.ExprOp)):
with T.If(self.eval_expr(node.test)):
if not predicate.dtype.is_bool():
self.report_error(node.test, f"If condition must be a boolean expression, but got {predicate.dtype}")
with T.If(predicate):
with T.Then():
with self.var_table.with_frame():
self.visit_body(node.body)
Expand Down Expand Up @@ -519,6 +521,8 @@ def visit_assert(self: Parser, node: doc.Assert) -> None:
"""
cond = self.eval_expr(node.test)
msg = self.eval_expr(node.msg)
if isinstance(cond, (PrimExpr, tvm.tirx.expr.ExprOp)) and not cond.dtype.is_bool():
self.report_error(node.test, f"Assert condition must be a boolean expression, but got {cond.dtype}")
frame = T.Assert(cond, msg)
frame.add_callback(partial(frame.__exit__, None, None, None))
frame.__enter__()
Expand Down
Loading