This is an automated email from the ASF dual-hosted git repository.

tqchen pushed a commit to branch tvmscript-ast-only-transpiler
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit 063ba1330b28faf9beda7a2a7a49820c9ff9da4c
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 17:50:15 2026 +0000

    Build keyword expressions eagerly with conditional runtime evaluation
---
 python/tvm/relax/script/builder/__init__.py        |  99 +++++++++
 python/tvm/s_tir/tensor_intrin/rocm.py             |  11 +-
 python/tvm/script/ir_builder/protocol.py           | 130 ++++++-----
 python/tvm/script/parser/transpile.py              |  83 +++----
 python/tvm/tirx/script/builder/__init__.py         |  73 ++++++
 src/tirx/script/printer/expr.cc                    |   7 +-
 src/tirx/transform/lower_tvm_builtin.cc            |  59 ++++-
 .../python/tvmscript/test_builder_short_circuit.py | 246 +++++++++++++++++++++
 tests/python/tvmscript/test_parser.py              |  52 +++++
 .../python/tvmscript/test_tvmscript_printer_tir.py |   4 +-
 10 files changed, 663 insertions(+), 101 deletions(-)

diff --git a/python/tvm/relax/script/builder/__init__.py 
b/python/tvm/relax/script/builder/__init__.py
index 30e845fd88..3a5b83058e 100644
--- a/python/tvm/relax/script/builder/__init__.py
+++ b/python/tvm/relax/script/builder/__init__.py
@@ -532,3 +532,102 @@ _protocol.register_call_kind(_sys.modules[__name__], 
_ir.GlobalVar, _global_call
 
 
 for_ = For
+
+
+def if_then_else_(condition, true_value, false_value):
+    """Construct scalar control flow with eager Python and selected-arm 
runtime use.
+
+    A host condition selects its already-built value. Primitive branches use
+    the native conditional intrinsic; Relax values use If so normalization and
+    VM lowering keep each branch's calls inside that branch. Tensor conditions
+    must be scalar booleans, as required by Relax If. This operation does not
+    emit either branch into an outer block or enter a builder frame. Returns a
+    selected host value, primitive conditional expression, or Relax If. Branch
+    types must be compatible; invalid types/conditions raise errors during
+    construction or subsequent Relax normalization. No state is retained.
+    """
+    if isinstance(condition, _ffi.ObjectConvertible):
+        condition = condition.asobject()
+    if not isinstance(condition, _ir.Expr):
+        return true_value if condition else false_value
+    true_value = (
+        true_value.asobject() if isinstance(true_value, 
_ffi.ObjectConvertible) else true_value
+    )
+    false_value = (
+        false_value.asobject() if isinstance(false_value, 
_ffi.ObjectConvertible) else false_value
+    )
+    if _ir.is_prim_expr(condition) and all(
+        _ir.is_prim_expr(value)
+        if isinstance(value, _ir.Expr)
+        else isinstance(value, _numbers.Number)
+        for value in (true_value, false_value)
+    ):
+        return _tir.if_then_else(condition, true_value, false_value)
+    return _relax.If(condition, _value(true_value), _value(false_value))
+
+
+def _chain_binding(variable, value, body):
+    if _ir.is_prim_expr(value) and _ir.is_prim_expr(body):
+        return _tir.Let(variable, value, body)
+    return _relax.SeqExpr([_relax.BindingBlock([_relax.VarBinding(variable, 
value)])], body)
+
+
+def and_(*values, chain=None):
+    """Construct scalar conjunction with short-circuit execution in compiled 
code.
+
+    Python eagerly builds all values. ``chain`` is an optional tuple of the
+    original comparison operands; it creates progressively scoped expression
+    bindings so shared operands execute once at runtime, before their first
+    comparison. Neither branch nor a later operand is normalized/emitted in the
+    outer scope. No source callbacks or persistent construction state are used.
+    Inputs accept host values, primitive boolean expressions, and scalar 
boolean
+    tensors. Returns a host value, primitive expression, or Relax If/SeqExpr.
+    chain defaults to None and otherwise must contain len(values)+1 operands;
+    values must be their ordered comparisons. Empty values raise TypeError;
+    invalid chain length raises ValueError. Invalid tensor/type conditions fail
+    during construction or Relax normalization. No frame is entered.
+    """
+    if not values:
+        raise TypeError("and_ requires at least one operand")
+    if chain is not None:
+        return _protocol._comparison_chain(values, chain, and_, _chain_binding)
+    result = values[-1]
+    for value in reversed(values[:-1]):
+        result = if_then_else_(
+            value, result, False if isinstance(value, _ir.Expr | 
_ffi.ObjectConvertible) else value
+        )
+    return result
+
+
+def or_(*values):
+    """Construct scalar runtime disjunction from eagerly built Python 
arguments.
+
+    The right operand is evaluated only when the left is false. Conditions
+    follow Relax If's scalar boolean contract. Inputs accept host values,
+    primitive boolean expressions, or scalar boolean tensors. Returns a host
+    value, primitive expression, or Relax If. Empty values raise TypeError;
+    invalid types/conditions fail in construction or Relax normalization. No
+    outer bindings are emitted, frames entered, or persistent state retained.
+    """
+    if not values:
+        raise TypeError("or_ requires at least one operand")
+    result = values[-1]
+    for value in reversed(values[:-1]):
+        result = if_then_else_(
+            value, True if isinstance(value, _ir.Expr | 
_ffi.ObjectConvertible) else value, result
+        )
+    return result
+
+
+def not_(value):
+    """Return host, scalar primitive, or tensor boolean negation without 
coercion.
+
+    value accepts a host value or boolean primitive/tensor expression. Returns 
a
+    host bool, primitive Not, or Relax logical_not call. Invalid IR types fail
+    during construction/normalization; host truth-conversion errors propagate.
+    This operation emits no bindings, enters no frames, and retains no state.
+    """
+    return logical_not(value)
+
+
+__all__ += ["and_", "if_then_else_", "not_", "or_"]
diff --git a/python/tvm/s_tir/tensor_intrin/rocm.py 
b/python/tvm/s_tir/tensor_intrin/rocm.py
index 29749dd443..902df74b7a 100644
--- a/python/tvm/s_tir/tensor_intrin/rocm.py
+++ b/python/tvm/s_tir/tensor_intrin/rocm.py
@@ -258,6 +258,9 @@ def get_mfma_load_intrin(
 def get_mfma_intrin(k_dim, in_dtype="float32", out_dtype="float32", 
b_transposed=False):
     local_size = (M_DIM * k_dim) // WARP_SIZE
     local_size_out = (M_DIM * N_DIM) // WARP_SIZE
+    # This is a host configuration choice: a one-lane Ramp is invalid to
+    # construct. Resolve it before DSL expressions eagerly construct both arms.
+    local_index = T.ramp(0, 1, local_size) if local_size > 1 else 0
     if k_dim == 4:
         index_map_A = shared_16x4_to_local_64x1_layout_A
         index_map_B = shared_4x16_to_local_64x1_layout_B
@@ -336,8 +339,8 @@ def get_mfma_intrin(k_dim, in_dtype="float32", 
out_dtype="float32", b_transposed
             T.launch_thread(tx, WARP_SIZE)
             C[tx, T.ramp(0, 1, local_size_out)] = T.call_llvm_pure_intrin(
                 T.llvm_lookup_intrinsic_id(mfma_intrin),
-                A[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0],
-                B[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0],
+                A[tx, local_index],
+                B[tx, local_index],
                 C[tx, T.ramp(0, 1, local_size_out)],
                 T.int32(0),
                 T.int32(0),
@@ -366,12 +369,12 @@ def get_mfma_intrin(k_dim, in_dtype="float32", 
out_dtype="float32", b_transposed
                 T.call_intrin(
                     "int32",
                     "tirx.reinterpret",
-                    A[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0],
+                    A[tx, local_index],
                 ),
                 T.call_intrin(
                     "int32",
                     "tirx.reinterpret",
-                    B[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0],
+                    B[tx, local_index],
                 ),
                 C[tx, T.ramp(0, 1, local_size_out)],
                 T.int32(0),
diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index 5d09a7a683..ff70228b19 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -230,59 +230,69 @@ def is_python_bool(value):
     return isinstance(value, bool)
 
 
-def compare_chain(logical_and, operands, comparisons):
-    """Evaluate a comparison chain in source order with host short-circuiting.
-
-    logical_and is the builder binary conjunction callable; operands is a
-    nonempty sequence of zero-argument callbacks; comparisons has one binary
-    callback per adjacent operand pair. Returns a host bool or builder result.
-    Each operand runs at most once. A Python False stops further evaluation;
-    symbolic conjunction uses logical_and. Callback/arity errors propagate.
-    No frames or persistent state are created by this helper.
+def _comparison_chain(comparisons, operands, conjunction, bind):
+    """Bind original chain operands progressively around short-circuit tests.
+
+    All inputs have already been constructed eagerly. The operand tuple retains
+    pre-conversion expression identities: comparison overloads may cast one
+    middle operand differently on either side. Python object reuse alone does
+    not make its compiled value execute once. Fresh IR variables replace those
+    identities inside comparisons, and each binding occurs only when its
+    comparison is reached. The first two operands evaluate left to right.
+
+    conjunction and bind are builder implementation functions, not lazy source
+    callbacks. They construct control flow and a dialect-specific expression
+    binding. All maps/variables are local to this construction; no registry,
+    native frame, or source expression is mutated.
     """
-    left = operands[0]()
-    result = True
-    for index, comparison in enumerate(comparisons):
-        right = operands[index + 1]()
-        current = comparison(left, right)
-        result = current if index == 0 else logical_and(result, current)
-        if isinstance(result, bool) and not result:
-            return False
-        left = right
-    return result
-
-
-def logical_chain(operation, operands, short_circuit):
-    """Fold lazy operands with host boolean short-circuiting.
-
-    operation is the builder binary and/or callable; operands is a nonempty
-    sequence of zero-argument callbacks. short_circuit is False for and, True
-    for or. Returns the first short-circuit bool or final builder result.
-    Callback errors propagate. Only evaluation-local state is retained and
-    no frames are entered; symbolic values never undergo Python truth tests.
-    """
-    result = operands[0]()
-    for operand in operands[1:]:
-        if isinstance(result, bool) and result is short_circuit:
-            return result
-        result = operation(result, operand())
+    import tvm_ffi
+
+    if len(operands) != len(comparisons) + 1:
+        raise ValueError("A comparison chain requires one more operand than 
comparison")
+    replacements = []
+    bindings = []
+    for operand in operands:
+        if isinstance(operand, tvm_ffi.ObjectConvertible):
+            operand = operand.asobject()
+        if not isinstance(operand, ir.Expr) or isinstance(operand, ir.Var):
+            bindings.append(None)
+            continue
+        previous = next((var for value, var in replacements if 
value.same_as(operand)), None)
+        if previous is not None:
+            bindings.append(None)
+            continue
+        variable = ir.Var("chain_operand", operand.ty)
+        replacements.append((operand, variable))
+        bindings.append((variable, operand))
+
+    def replace(value, mutator):
+        for original, variable in replacements:
+            if value.same_as(original):
+                return variable
+        return mutator.default_mutate(value)
+
+    conditions = []
+    for comparison in comparisons:
+        if isinstance(comparison, tvm_ffi.ObjectConvertible):
+            comparison = comparison.asobject()
+        conditions.append(
+            tvm_ffi.structural_mutate(comparison, [(ir.Expr, replace)])
+            if isinstance(comparison, ir.Expr)
+            else comparison
+        )
+    result = conditions[-1]
+    for index in range(len(conditions) - 1, -1, -1):
+        if index < len(conditions) - 1:
+            result = conjunction(conditions[index], result)
+        if bindings[index + 1] is not None:
+            variable, value = bindings[index + 1]
+            result = bind(variable, value, result)
+    if bindings[0] is not None:
+        variable, value = bindings[0]
+        result = bind(variable, value, result)
     return result
 
 
-def select_lazy(operation, condition, true_value, false_value):
-    """Evaluate a conditional expression through lazy source callbacks.
-
-    operation accepts (condition, true_result, false_result); condition is an
-    already evaluated host/builder value. true_value and false_value are
-    zero-argument callbacks. A Python bool evaluates and returns one arm;
-    otherwise both arms run in source order and operation returns the result.
-    Callback errors propagate; this helper owns no frames or persistent state.
-    """
-    if isinstance(condition, bool):
-        return true_value() if condition else false_value()
-    return operation(condition, true_value(), false_value())
-
-
 def register_call_kind(builder, value_type, adapter):
     """Register runtime adaptation of callable values for one builder 
namespace.
 
@@ -299,18 +309,30 @@ def register_call_kind(builder, value_type, adapter):
     builder.__tvm_call_kinds__ = policies
 
 
-def callee(builder, value):
+def callee(builder, value, *, span=None):
     """Adapt a runtime callable according to one builder's registered policy.
 
     ``builder`` is the construction namespace and ``value`` the already
     evaluated call target. Returns the first matching adapter's result, or
-    value unchanged when no type matches. The operation reads namespace-owned
-    registration state, enters no frames itself, and propagates adapter errors.
+    value unchanged when no type matches. Optional span accepts source_span
+    input forms and wraps the actual construction call in the private builder
+    span context, including operations which emit statements and return None.
+    None adds no wrapper or span instrumentation. Registration is read-only;
+    adapter and call errors propagate with original types and source metadata.
     """
     for value_type, adapter in getattr(builder, "__tvm_call_kinds__", 
{}).items():
         if isinstance(value, value_type):
-            return adapter(value)
-    return value
+            value = adapter(value)
+            break
+    if span is None:
+        return value
+
+    @wraps(value)
+    def construct(*args, **kwargs):
+        with _construction_span(span):
+            return value(*args, **kwargs)
+
+    return construct
 
 
 def __getattr__(name):
diff --git a/python/tvm/script/parser/transpile.py 
b/python/tvm/script/parser/transpile.py
index 6f45b83e5e..e78d9fbec7 100644
--- a/python/tvm/script/parser/transpile.py
+++ b/python/tvm/script/parser/transpile.py
@@ -315,39 +315,34 @@ class IRBuilderTranspiler(ast.NodeTransformer):
             isinstance(node, ast.NamedExpr) and not getattr(node, 
"_tvm_signature_binding", False)
         ):
             self._error(original, f"Unsupported expression: 
{type(node).__name__}")
-        # Pattern: not e -> X.logical_not(e).
+        # Pattern: not e -> X.not_(e). These operations construct IR eagerly;
+        # the builders, not Python truth testing, own compiled 
short-circuiting.
         if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not):
-            node = self._call(
-                self.dialect_prefix, "logical_not", 
[self._expression(node.operand)], node
-            )
-        # Pattern: a and/or b -> lazy host/builder logical_chain callbacks.
+            node = self._call(self.dialect_prefix, "not_", 
[self._expression(node.operand)], node)
+        # Pattern: a and/or b -> X.and_(a, b) / X.or_(a, b). Python evaluates
+        # every operand left-to-right while building; compiled IR is 
conditional.
         elif isinstance(node, ast.BoolOp):
-            method = "logical_and" if isinstance(node.op, ast.And) else 
"logical_or"
-            operands = [self._lambda([], self._expression(value), value) for 
value in node.values]
+            method = "and_" if isinstance(node.op, ast.And) else "or_"
             node = self._call(
-                self.infrastructure_name,
-                "logical_chain",
-                [
-                    self._attribute(self.dialect_prefix, method, node),
-                    ast.Tuple(operands, ast.Load()),
-                    ast.Constant(isinstance(node.op, ast.Or)),
-                ],
+                self.dialect_prefix,
+                method,
+                [self._expression(value) for value in node.values],
                 node,
             )
-        # Pattern: a if c else b -> lazy select callbacks, preserving host 
effects.
+        # Pattern: a if c else b -> X.if_then_else_(c, a, b). Both arms are
+        # constructed now; only the selected arm executes in compiled code.
         elif isinstance(node, ast.IfExp):
             node = self._call(
-                self.infrastructure_name,
-                "select_lazy",
+                self.dialect_prefix,
+                "if_then_else_",
                 [
-                    self._attribute(self.dialect_prefix, "select", node),
                     self._expression(node.test),
-                    self._lambda([], self._expression(node.body), node.body),
-                    self._lambda([], self._expression(node.orelse), 
node.orelse),
+                    self._expression(node.body),
+                    self._expression(node.orelse),
                 ],
                 node,
             )
-        # Pattern: a < b < c -> compare_chain, evaluating each operand once.
+        # Pattern: a < b < c -> X.and_(a < b, b < c, chain=(a, b, c)).
         elif isinstance(node, ast.Compare) and len(node.ops) > 1:
             node = self._compare_chain(node)
         # Pattern: f"text{e}" -> preserve fragments and translate only 
expressions.
@@ -360,12 +355,17 @@ class IRBuilderTranspiler(ast.NodeTransformer):
                         child.format_spec = 
self._format_spec(child.format_spec)
         else:
             self._expression_children(node)
+        # Pattern: f(args) -> at(loc, callee(X, f, span=loc)(args)). The
+        # builder adapter owns construction contexts for calls which emit a
+        # statement and return None, including inline helpers. No parser-side
+        # context manager or deferred source-expression callback is introduced.
         if isinstance(node, ast.Call):
             node.func = self._call(
                 self.infrastructure_name,
                 "callee",
                 [self._name(self.dialect_prefix, original), node.func],
                 original,
+                span=self.span(original),
             )
         if isinstance(node, ast.Starred) or (
             isinstance(node, ast.Name) and isinstance(node.ctx, ast.Store)
@@ -388,25 +388,32 @@ class IRBuilderTranspiler(ast.NodeTransformer):
         return self._located(ast.Lambda(arguments, value), original)
 
     def _compare_chain(self, node):
-        operands = [
-            self._lambda([], self._expression(value), value)
-            for value in [node.left, *node.comparators]
-        ]
-        comparisons = []
-        for operation in node.ops:
-            left, right = self.fresh("left"), self.fresh("right")
-            comparison = self._located(
+        # An immediately invoked Python lambda binds eagerly evaluated operands
+        # once in source order. It passes ordinary overloaded comparisons to 
the
+        # builder, never lazy callbacks. chain=(a, b, c) records syntax 
provenance so
+        # builders also bind shared middle operands once at compiled runtime,
+        # where duplicating an effectful expression would change semantics.
+        operands = [node.left, *node.comparators]
+        names = [self.fresh("operand") for _ in operands]
+        comparisons = [
+            self._located(
                 ast.Compare(self._name(left, node), [operation], 
[self._name(right, node)]), node
             )
-            comparisons.append(self._lambda([left, right], comparison, node))
-        return self._call(
-            self.infrastructure_name,
-            "compare_chain",
-            [
-                self._attribute(self.dialect_prefix, "logical_and", node),
-                ast.Tuple(operands, ast.Load()),
-                ast.Tuple(comparisons, ast.Load()),
-            ],
+            for left, operation, right in zip(names, node.ops, names[1:])
+        ]
+        value = self._call(
+            self.dialect_prefix,
+            "and_",
+            comparisons,
+            node,
+            chain=ast.Tuple([self._name(name, node) for name in names], 
ast.Load()),
+        )
+        return self._located(
+            ast.Call(
+                self._lambda(names, value, node),
+                [self._expression(value) for value in operands],
+                [],
+            ),
             node,
         )
 
diff --git a/python/tvm/tirx/script/builder/__init__.py 
b/python/tvm/tirx/script/builder/__init__.py
index 27da1890dc..10537d2d08 100644
--- a/python/tvm/tirx/script/builder/__init__.py
+++ b/python/tvm/tirx/script/builder/__init__.py
@@ -779,3 +779,76 @@ def _global_callee(function):
 
 
 _register_call_kind(_sys.modules[__name__], _ir.GlobalVar, _global_callee)
+
+
+def if_then_else_(condition, true_value, false_value):
+    """Construct a scalar conditional whose compiled code evaluates one arm.
+
+    Arguments are already constructed eagerly by Python. A host condition
+    selects a host value; a scalar IR condition emits the native control-flow
+    intrinsic, not an eager Select. Neither arm is emitted as an outer binding.
+    Returns the selected host value or a primitive IR expression. condition may
+    also be an ObjectConvertible scalar expression. IR arms must have 
compatible
+    primitive types; invalid conditions/types raise native TypeError or
+    ValueError. This operation owns no frame/state.
+    """
+    return select(condition, true_value, false_value)
+
+
+def and_(*values, chain=None):
+    """Construct a scalar conjunction with left-to-right runtime short circuit.
+
+    Python eagerly constructs every argument. Runtime evaluation stops at the
+    first false value. ``chain`` optionally carries the original, already-built
+    comparison operands, allowing native Let bindings to evaluate each shared
+    middle operand once, after the preceding comparison succeeds. This syntax
+    metadata does not contain callbacks. Host values preserve Python selection;
+    vector conditions are rejected by the scalar conditional intrinsic.
+    Returns a selected host value or scalar boolean IR expression. chain 
defaults
+    to None; when supplied it must have len(values)+1 original operands, with
+    values being their ordered comparisons. A malformed length raises 
ValueError;
+    empty values raise TypeError and invalid IR types propagate native errors.
+    No construction frames or state persist after the call.
+    """
+    if not values:
+        raise TypeError("and_ requires at least one operand")
+    if chain is not None:
+        from tvm.script.ir_builder.protocol import _comparison_chain
+
+        return _comparison_chain(values, chain, and_, _tir.Let)
+    result = values[-1]
+    for value in reversed(values[:-1]):
+        result = if_then_else_(
+            value, result, False if isinstance(value, _ir.Expr | 
_ffi.ObjectConvertible) else value
+        )
+    return result
+
+
+def or_(*values):
+    """Construct eager arguments into a scalar runtime short-circuit 
disjunction.
+
+    Runtime evaluation stops at the first true operand. Python has already
+    constructed all arguments. Host values preserve Python selection; native
+    scalar-condition/type checks apply to IR operands. Returns a selected host
+    value or scalar boolean IR expression. Empty values raise TypeError; 
invalid
+    IR operand types propagate native errors. No frame is entered or state 
kept.
+    """
+    if not values:
+        raise TypeError("or_ requires at least one operand")
+    result = values[-1]
+    for value in reversed(values[:-1]):
+        result = if_then_else_(
+            value, True if isinstance(value, _ir.Expr | 
_ffi.ObjectConvertible) else value, result
+        )
+    return result
+
+
+def not_(value):
+    """Return host or IR boolean negation without testing an IR value in 
Python.
+
+    value accepts a host value, primitive boolean expression, or 
ObjectConvertible
+    expression. Returns a host bool or primitive boolean expression. 
Unsupported
+    IR types propagate native errors; host truth-conversion errors propagate.
+    This operation enters no frame and retains no state.
+    """
+    return logical_not(value)
diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc
index f694b844ed..3b992191b6 100644
--- a/src/tirx/script/printer/expr.cc
+++ b/src/tirx/script/printer/expr.cc
@@ -582,8 +582,11 @@ TVM_FFI_STATIC_INIT_BLOCK() {
   TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(NE, prim::NENode, not_equal, "NE", 
kNotEq);
   TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(GT, prim::GTNode, greater, "GT", 
kGt);
   TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(GE, prim::GENode, greater_equal, 
"GE", kGtE);
-  TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(And, prim::AndNode, logical_and, 
"And", kAnd);
-  TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(Or, prim::OrNode, logical_or, "Or", 
kOr);
+  // Python keyword expressions construct short-circuit control-flow IR. These
+  // native eager boolean nodes need explicit constructors to preserve their
+  // semantics and structure when the printed program is parsed again.
+  TVM_SCRIPT_PRINTER_DEF_BINARY(And, "And");
+  TVM_SCRIPT_PRINTER_DEF_BINARY(Or, "Or");
 
   TVM_SCRIPT_PRINTER_DEF_BINARY(Mod, "truncmod");
   TVM_SCRIPT_PRINTER_DEF_BINARY(Min, "min");
diff --git a/src/tirx/transform/lower_tvm_builtin.cc 
b/src/tirx/transform/lower_tvm_builtin.cc
index 9f14fe4d95..c9ab00c647 100644
--- a/src/tirx/transform/lower_tvm_builtin.cc
+++ b/src/tirx/transform/lower_tvm_builtin.cc
@@ -34,6 +34,7 @@
 #include <tvm/tirx/stmt_functor.h>
 #include <tvm/tirx/transform.h>
 
+#include <deque>
 #include <unordered_set>
 
 #include "ir_utils.h"
@@ -409,7 +410,60 @@ class BuiltinLower : public StmtExprMutator {
     return IfThenElse(condition, then_case, else_case, op->span);
   }
 
+  UnchangedOr<PrimExpr> Mutate_(const prim::LetNode* op, InplaceMode 
inplace_mode) final {
+    // A packed call's argument preparation must precede its evaluation, not 
the
+    // entire enclosing expression. Realize each expression binding before
+    // visiting its body, so sibling calls cannot overwrite its argument slots.
+    Expr value = this->Mutate(op->value, 
inplace_mode).ValueOrUnchanged(op->value);
+    prep_seq_stack_.back().push_back(Bind(op->var, value, op->span));
+    return this->Mutate(op->body, inplace_mode).ValueOrUnchanged(op->body);
+  }
+
+  Expr MakeConditional(const CallNode* op, InplaceMode inplace_mode) {
+    Expr condition = this->Mutate(op->args[0], 
inplace_mode).ValueOrUnchanged(op->args[0]);
+    // Branch-local setup includes packed argument stores and bindings. Moving
+    // it outside the conditional both executes unselected operands and lets 
one
+    // arm overwrite the other's reused packed-call argument stack.
+    auto mutate_branch = [&](const Expr& branch) {
+      prep_seq_stack_.emplace_back();
+      Expr value = this->Mutate(branch, inplace_mode).ValueOrUnchanged(branch);
+      auto preparation = std::move(prep_seq_stack_.back());
+      prep_seq_stack_.pop_back();
+      return std::make_pair(std::move(value), std::move(preparation));
+    };
+    auto [true_value, true_prep] = mutate_branch(op->args[1]);
+    auto [false_value, false_prep] = mutate_branch(op->args[2]);
+    if (true_prep.empty() && false_prep.empty()) {
+      return Call(op->ty, op->op, {condition, true_value, false_value}, 
op->attrs, op->ty_args,
+                  op->span);
+    }
+
+    // A local scalar slot carries the selected result across statement 
branches.
+    // Native codegen promotes this slot to SSA where possible. Pointer-valued
+    // conditionals use an integer slot and preserve their bits by reinterpret.
+    auto primitive_type = op->ty.as<PrimType>();
+    TVM_FFI_ICHECK(primitive_type.has_value() || op->ty.as<PointerTypeNode>())
+        << "A conditional result must have primitive or pointer type";
+    PrimType storage_type = primitive_type.value_or(PrimType::UInt(64));
+    BufferVar result = decl_buffer({IntImm::Int32(1)}, storage_type, 
"conditional_result", "local");
+    auto store_value = [&](const Expr& value) {
+      return primitive_type.has_value() ? value.as_or_throw<PrimExpr>()
+                                        : reinterpret(storage_type, 
value).as_or_throw<PrimExpr>();
+    };
+    true_prep.push_back(BufferStore(result, store_value(true_value), 
{IntImm::Int32(0)}));
+    false_prep.push_back(BufferStore(result, store_value(false_value), 
{IntImm::Int32(0)}));
+    auto& preparation = prep_seq_stack_.back();
+    preparation.push_back(AllocBuffer(result));
+    preparation.push_back(IfThenElse(condition.as_or_throw<PrimExpr>(), 
SeqStmt::Flatten(true_prep),
+                                     SeqStmt::Flatten(false_prep), op->span));
+    Expr loaded = BufferLoad(result, {IntImm::Int32(0)});
+    return primitive_type.has_value() ? loaded : reinterpret(op->ty, loaded);
+  }
+
   UnchangedOr<Expr> Mutate_(const CallNode* op, InplaceMode inplace_mode) 
final {
+    if (op->op.same_as(prim::builtin::if_then_else())) {
+      return MakeConditional(op, inplace_mode);
+    }
     if (op->op.same_as(builtin::tensormap_encode_tiled()) && 
!preserve_ffi_kernel_) {
       const auto* attr = op->attrs.as<TensorMapEncodeTiledAttr>();
       TVM_FFI_CHECK(attr && attr->rank >= 1 && attr->rank <= 5 &&
@@ -818,7 +872,10 @@ class BuiltinLower : public StmtExprMutator {
   }
 
   // The prepration sequence to be emitted before the current statement.
-  std::vector<std::vector<Stmt>> prep_seq_stack_;
+  // Packed-call lowering keeps a reference to its preparation list while
+  // recursively visiting conditional arguments. Branch scopes must not
+  // invalidate that reference when pushing another preparation list.
+  std::deque<std::vector<Stmt>> prep_seq_stack_;
   ffi::Optional<PrimExpr> device_type_{std::nullopt};
   ffi::Optional<PrimExpr> device_id_{std::nullopt};
 
diff --git a/tests/python/tvmscript/test_builder_short_circuit.py 
b/tests/python/tvmscript/test_builder_short_circuit.py
new file mode 100644
index 0000000000..5c84459ab5
--- /dev/null
+++ b/tests/python/tvmscript/test_builder_short_circuit.py
@@ -0,0 +1,246 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""Eager builder construction retains lazy execution of scalar IR 
expressions."""
+
+import numpy as np
+import pytest
+import tvm_ffi
+
+import tvm
+import tvm.testing
+from tvm import ir, relax, tirx
+from tvm.relax.script import builder as R
+from tvm.script import tirx as ST
+from tvm.script.ir_builder import IRBuilder
+from tvm.testing import env
+from tvm.tirx.script import builder as T
+
+pytestmark = pytest.mark.skipif(not env.has_llvm(), reason="LLVM support 
required")
+
+
[email protected]
+def effects():
+    calls = []
+
+    @tvm_ffi.register_global_func("test.builder.short_circuit", override=True)
+    def effect(label, value):
+        calls.append(label)
+        if label < 0:
+            raise RuntimeError("unselected expression executed")
+        return value
+
+    @tvm_ffi.register_global_func("test.builder.short_circuit_tensor", 
override=True)
+    def tensor_effect(label, value):
+        return tvm.runtime.tensor(np.array(effect(label, value), dtype="bool"))
+
+    return calls
+
+
+def _primitive_function(expression):
+    with IRBuilder() as builder:
+        with T.prim_func():
+            T.func_name("main")
+            flag = T.arg("flag", T.int32())
+            output = T.arg("output", T.Buffer((1,), "int32"))
+            value = expression(flag)
+            T.buffer_store(output, tirx.Cast("int32", value), [0])
+    return tvm.compile(builder.get(), target="llvm")
+
+
+def _run_primitive(function, flag):
+    output = tvm.runtime.tensor(np.zeros(1, dtype="int32"))
+    function(flag, output)
+    return int(output.numpy()[0])
+
+
+def _effect(label, value):
+    return tirx.call_packed("test.builder.short_circuit", label, value)
+
+
[email protected]("dialect", [T, R], ids=["tir", "relax-primitive"])
[email protected]("operation", ["and_", "or_", "if_then_else_"])
+def test_primitive_short_circuit(dialect, operation, effects):
+    constructed = []
+
+    def build(flag):
+        def operand(label, value):
+            constructed.append(label)
+            return _effect(label, value) != 0
+
+        if operation == "if_then_else_":
+            return dialect.if_then_else_(flag > 0, operand(1, 1), operand(2, 
0))
+        return getattr(dialect, operation)(flag > 0, operand(1, 1))
+
+    function = _primitive_function(build)
+    assert constructed == ([1, 2] if operation == "if_then_else_" else [1])
+    for flag in (0, 1):
+        effects.clear()
+        result = _run_primitive(function, flag)
+        if operation == "and_":
+            assert result == flag
+            assert effects == ([1] if flag else [])
+        elif operation == "or_":
+            assert result == 1
+            assert effects == ([] if flag else [1])
+        else:
+            assert result == flag
+            assert effects == ([1] if flag else [2])
+
+
[email protected]("operation", ["and_", "or_", "if_then_else_"])
+def test_unselected_invalid_expression(operation, effects):
+    def build(flag):
+        invalid = _effect(-1, 1) != 0
+        if operation == "if_then_else_":
+            return T.if_then_else_(flag > 0, True, invalid)
+        return getattr(T, operation)(flag > 0, invalid)
+
+    function = _primitive_function(build)
+    flag = 0 if operation == "and_" else 1
+    assert _run_primitive(function, flag) == flag
+    assert effects == []
+    with pytest.raises(RuntimeError, match="unselected expression executed"):
+        _run_primitive(function, 1 - flag)
+    assert effects == [-1]
+
+
[email protected]("dialect", [T, R], ids=["tir", "relax-primitive"])
+def test_chained_comparison_runtime_single_evaluation(dialect, effects):
+    def build(flag):
+        operands = (_effect(0, 0), _effect(1, flag), _effect(2, 3), _effect(3, 
4))
+        a, b, c, d = operands
+        return dialect.and_(a < b, b < c, c < d, chain=operands)
+
+    function = _primitive_function(build)
+    for flag, result, expected_calls in [(-1, 0, [0, 1]), (2, 1, [0, 1, 2, 
3]), (5, 0, [0, 1, 2])]:
+        effects.clear()
+        assert _run_primitive(function, flag) == result
+        assert effects == expected_calls
+
+
[email protected]("operation", ["and_", "or_", "if_then_else_"])
+def test_relax_vm_short_circuit(operation, effects):
+    flag = ir.Var("flag", relax.TensorType([], "bool"))
+
+    def operand(label, value):
+        return R.call_packed(
+            "test.builder.short_circuit_tensor",
+            label,
+            value,
+            ty_args=relax.TensorType([], "bool"),
+        )
+
+    if operation == "if_then_else_":
+        expression = R.if_then_else_(flag, operand(1, 1), operand(2, 0))
+    else:
+        expression = getattr(R, operation)(flag, operand(1, 1))
+    builder = relax.BlockBuilder()
+    with builder.function("main", [flag], pure=False):
+        builder.emit_func_output(builder.emit(expression))
+    executable = relax.build(builder.get(), target="llvm")
+    vm = relax.VirtualMachine(executable, tvm.cpu())
+    for value in (False, True):
+        effects.clear()
+        result = bool(vm["main"](tvm.runtime.tensor(np.array(value))).numpy())
+        if operation == "and_":
+            assert result == value
+            assert effects == ([1] if value else [])
+        elif operation == "or_":
+            assert result
+            assert effects == ([] if value else [1])
+        else:
+            assert result == value
+            assert effects == ([1] if value else [2])
+
+
+def test_source_chain_execution(effects):
+    @ST.prim_func
+    def compute(flag: ST.int32, output: ST.Buffer((1,), "int32")):
+        output[0] = ST.Cast(
+            "int32", _effect(0, 0) < _effect(1, flag) < _effect(2, 3) < 
_effect(3, 4)
+        )
+
+    function = tvm.compile(compute, target="llvm")
+    for flag, result, calls in [(-1, 0, [0, 1]), (2, 1, [0, 1, 2, 3])]:
+        effects.clear()
+        assert _run_primitive(function, flag) == result
+        assert effects == calls
+
+
+def test_comparison_casts_preserve_original_middle(effects):
+    def build(flag):
+        left = tirx.const(16777216, "float32")
+        middle = _effect(1, flag)
+        right = tirx.const(16777216, "int32")
+        # The left comparison rounds the middle to float32. The right 
comparison
+        # must use the original int32 value rather than that rounded result.
+        return T.and_(left <= middle, middle > right, chain=(left, middle, 
right))
+
+    function = _primitive_function(build)
+    assert _run_primitive(function, 16777217) == 1
+    assert effects == [1]
+
+
+def test_relax_vm_chain_execution(effects):
+    @tvm_ffi.register_global_func("test.builder.chain_tensor", override=True)
+    def effect(label, value):
+        effects.append(label)
+        return tvm.runtime.tensor(np.array(value, dtype="int32"))
+
+    flag = ir.Var("flag", ir.PrimType("int32"))
+
+    def operand(label, value):
+        return R.call_packed(
+            "test.builder.chain_tensor",
+            label,
+            value,
+            ty_args=relax.TensorType([], "int32"),
+        )
+
+    operands = (operand(0, 0), operand(1, flag), operand(2, 3), operand(3, 4))
+    a, b, c, d = operands
+    expression = R.and_(a < b, b < c, c < d, chain=operands)
+    builder = relax.BlockBuilder()
+    with builder.function("main", [flag], pure=False):
+        builder.emit_func_output(builder.emit(expression))
+    vm = relax.VirtualMachine(relax.build(builder.get(), target="llvm"), 
tvm.cpu())
+    for value, result, calls in [
+        (-1, False, [0, 1]),
+        (2, True, [0, 1, 2, 3]),
+        (5, False, [0, 1, 2]),
+    ]:
+        effects.clear()
+        assert bool(vm["main"](value).numpy()) == result
+        assert effects == calls
+
+
+def test_conditional_inside_packed_argument(effects):
+    # Lowering the outer call retains its preparation list while recursively
+    # lowering branch-local setup. Both stack slots and list references must
+    # remain valid across that nested scope.
+    def build(flag):
+        return _effect(9, T.if_then_else_(flag > 0, _effect(1, 11), _effect(2, 
22)))
+
+    function = _primitive_function(build)
+    for flag, result, calls in [(0, 22, [2, 9]), (1, 11, [1, 9])]:
+        effects.clear()
+        assert _run_primitive(function, flag) == result
+        assert effects == calls
+
+
+if __name__ == "__main__":
+    tvm.testing.main()
diff --git a/tests/python/tvmscript/test_parser.py 
b/tests/python/tvmscript/test_parser.py
index 55d6f1877c..a029cadc52 100644
--- a/tests/python/tvmscript/test_parser.py
+++ b/tests/python/tvmscript/test_parser.py
@@ -603,3 +603,55 @@ def test_optional_shared_source_metadata(track_span, 
monkeypatch):
         )
     assert "<tracking-error>:" in str(error.value)
     assert "missing_name" in str(error.value)
+
+
+def test_keyword_expressions_construct_all_operands_in_order():
+    events, chains = [], []
+
+    class Builder(_Recorder):
+        def and_(self, *values, chain=None):
+            if chain is not None:
+                chains.append(chain)
+            return all(values)
+
+        def or_(self, *values):
+            return any(values)
+
+        def not_(self, value):
+            return not value
+
+        def if_then_else_(self, condition, true_value, false_value):
+            return true_value if condition else false_value
+
+    def value(label, result):
+        events.append(label)
+        return result
+
+    builder = Builder()
+    compiler = Compiler(
+        """
[email protected]
+def f():
+    a = value("and-left", False) and value("and-right", True)
+    b = value("or-left", True) or value("or-right", False)
+    c = value("true", 10) if value("condition", True) else value("false", 20)
+    d = value("chain-left", 1) < value("chain-middle", 2) < 
value("chain-right", 3)
+    return (a, b, c, d, not a)
+""",
+        {"D": SimpleNamespace(function=make_decorator(builder)), "value": 
value},
+    )
+    _run(compiler, compiler.tree.body[0], builder)
+    assert events == [
+        "and-left",
+        "and-right",
+        "or-left",
+        "or-right",
+        "condition",
+        "true",
+        "false",
+        "chain-left",
+        "chain-middle",
+        "chain-right",
+    ]
+    assert chains == [(1, 2, 3)]
+    assert builder.returned == [(False, True, 10, True, True)]
diff --git a/tests/python/tvmscript/test_tvmscript_printer_tir.py 
b/tests/python/tvmscript/test_tvmscript_printer_tir.py
index b1a911797d..434e1231dc 100644
--- a/tests/python/tvmscript/test_tvmscript_printer_tir.py
+++ b/tests/python/tvmscript/test_tvmscript_printer_tir.py
@@ -631,7 +631,7 @@ def test_logical():
         """
 a = T.bool()
 b = T.bool()
-a and b
+T.And(a, b)
 """,
     )
     _assert_print(
@@ -639,7 +639,7 @@ a and b
         """
 a = T.bool()
 b = T.bool()
-a or b
+T.Or(a, b)
 """,
     )
     _assert_print(

Reply via email to