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(
