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

tqchen pushed a commit to branch tvmscript-generic-parser-builder
in repository https://gitbox.apache.org/repos/asf/tvm.git

commit dd62dde13588ad211534909b5c76cf8db1af170f
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 02:42:53 2026 +0000

    Preserve lazy Python control flow during builder translation
---
 python/tvm/script/ir_builder/protocol.py      |  36 ++++++
 python/tvm/script/parser_v2/frontend.py       |   7 +-
 python/tvm/script/parser_v2/transform.py      | 169 ++++++++++++++++++++++----
 python/tvm/tirx/script/builder_v2/__init__.py |  12 +-
 4 files changed, 195 insertions(+), 29 deletions(-)

diff --git a/python/tvm/script/ir_builder/protocol.py 
b/python/tvm/script/ir_builder/protocol.py
index 1813bd61b7..740091eb2d 100644
--- a/python/tvm/script/ir_builder/protocol.py
+++ b/python/tvm/script/ir_builder/protocol.py
@@ -139,3 +139,39 @@ def require_defined(value, name):
     if value is MISSING:
         raise NameError(f"name {name!r} is not defined")
     return value
+
+
+def is_python_bool(value):
+    """Identify an ordinary condition that selects a construction-time 
branch."""
+    return isinstance(value, bool)
+
+
+def compare_chain(logical_and, operands, comparisons):
+    """Evaluate each operand once, stopping after an ordinary false 
comparison."""
+    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):
+    """Preserve ordinary boolean short-circuiting while constructing symbolic 
operands."""
+    result = operands[0]()
+    for operand in operands[1:]:
+        if isinstance(result, bool) and result is short_circuit:
+            return result
+        result = operation(result, operand())
+    return result
+
+
+def select_lazy(operation, condition, true_value, false_value):
+    """Select one ordinary boolean arm, or construct both symbolic arms."""
+    if isinstance(condition, bool):
+        return true_value() if condition else false_value()
+    return operation(condition, true_value(), false_value())
diff --git a/python/tvm/script/parser_v2/frontend.py 
b/python/tvm/script/parser_v2/frontend.py
index a54a407c28..e39bdc41a1 100644
--- a/python/tvm/script/parser_v2/frontend.py
+++ b/python/tvm/script/parser_v2/frontend.py
@@ -371,13 +371,8 @@ class Compiler:
                     node.func = ast.copy_location(ast.Name(name, ast.Load()), 
node.func)
             method = None
             values = []
-            if isinstance(node, ast.BoolOp):
-                method = "logical_and" if isinstance(node.op, ast.And) else 
"logical_or"
-                values = node.values
-            elif isinstance(node, ast.UnaryOp) and isinstance(node.op, 
ast.Not):
+            if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.Not):
                 method, values = "logical_not", [node.operand]
-            elif isinstance(node, ast.IfExp):
-                method, values = "select", [node.test, node.body, node.orelse]
             if method is not None:
                 node = ast.copy_location(
                     ast.Call(
diff --git a/python/tvm/script/parser_v2/transform.py 
b/python/tvm/script/parser_v2/transform.py
index 75ba52708f..80bd0fb1ec 100644
--- a/python/tvm/script/parser_v2/transform.py
+++ b/python/tvm/script/parser_v2/transform.py
@@ -178,7 +178,34 @@ class Transformer(ast.NodeTransformer):
             ast.fix_missing_locations(node)
         if isinstance(node, ast.Await | ast.Yield | ast.YieldFrom | 
ast.NamedExpr):
             self._error(original, f"Unsupported expression: 
{type(node).__name__}")
-        if isinstance(node, ast.JoinedStr):
+        if 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]
+            node = self._call(
+                self.infrastructure_name,
+                "logical_chain",
+                [
+                    self._attribute(self.builder_name, method, node),
+                    ast.Tuple(operands, ast.Load()),
+                    ast.Constant(isinstance(node.op, ast.Or)),
+                ],
+                node,
+            )
+        elif isinstance(node, ast.IfExp):
+            node = self._call(
+                self.infrastructure_name,
+                "select_lazy",
+                [
+                    self._attribute(self.builder_name, "select", node),
+                    self._expression(node.test),
+                    self._lambda([], self._expression(node.body), node.body),
+                    self._lambda([], self._expression(node.orelse), 
node.orelse),
+                ],
+                node,
+            )
+        elif isinstance(node, ast.Compare) and len(node.ops) > 1:
+            node = self._compare_chain(node)
+        elif isinstance(node, ast.JoinedStr):
             # JoinedStr's children must remain literal 
fragments/FormattedValue nodes.
             for child in node.values:
                 if isinstance(child, ast.FormattedValue):
@@ -195,6 +222,41 @@ class Transformer(ast.NodeTransformer):
             return node
         return self._call(self.infrastructure_name, "_at", 
[self.span(original), node], original)
 
+    def _lambda(self, names, value, original):
+        arguments = ast.arguments(
+            posonlyargs=[],
+            args=[ast.arg(arg=name) for name in names],
+            vararg=None,
+            kwonlyargs=[],
+            kw_defaults=[],
+            kwarg=None,
+            defaults=[],
+        )
+        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(
+                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.builder_name, "logical_and", node),
+                ast.Tuple(operands, ast.Load()),
+                ast.Tuple(comparisons, ast.Load()),
+            ],
+            node,
+        )
+
     def _format_spec(self, node):
         for child in node.values:
             if isinstance(child, ast.FormattedValue):
@@ -434,7 +496,7 @@ class Transformer(ast.NodeTransformer):
             visitor.visit(statement)
         return names
 
-    def _scope(self, body, original, prefix, initial=None):
+    def _scope(self, body, original, prefix, initial=None, 
return_bindings=False):
         outer_bound, outer_environment, outer_optional = self.bound, 
self.environment, self.optional
         referenced = {
             node.id
@@ -452,6 +514,23 @@ class Transformer(ast.NodeTransformer):
         }
         prefix_statements = [] if initial is None else initial()
         translated = prefix_statements + self.transform_statements(body)
+        if return_bindings and not self.preserve_return:
+            names = sorted(self._assigned_names(body))
+            values = [
+                self._name(name, original)
+                if name in self.bound
+                else copy.deepcopy(
+                    self.optional.get(
+                        name, self._attribute(self.infrastructure_name, 
"MISSING", original)
+                    )
+                )
+                for name in names
+            ]
+            translated.append(
+                self._located(
+                    ast.Return(ast.Dict([ast.Constant(name) for name in 
names], values)), original
+                )
+            )
         self.bound, self.environment, self.optional = outer_bound, 
outer_environment, outer_optional
         if self.preserve_return:
             return translated or [self._located(ast.Pass(), original)]
@@ -484,15 +563,20 @@ class Transformer(ast.NodeTransformer):
         invocation = self._located(ast.Call(self._name(helper, original), [], 
[]), original)
         return [definition, self._statement(invocation, original)]
 
-    def _exports(self, frame, candidates, original):
-        mapping_stmt, mapping = self._cache(
-            self._call(
-                self.infrastructure_name, "frame_result", [self._name(frame, 
original)], original
-            ),
-            original,
-            "exports",
-        )
-        result = [mapping_stmt]
+    def _exports(self, frame, candidates, original, mapping=None):
+        result = []
+        if mapping is None:
+            mapping_stmt, mapping = self._cache(
+                self._call(
+                    self.infrastructure_name,
+                    "frame_result",
+                    [self._name(frame, original)],
+                    original,
+                ),
+                original,
+                "exports",
+            )
+            result.append(mapping_stmt)
         for name in sorted(candidates):
             key = ast.Constant(name)
             condition = self._located(
@@ -525,24 +609,65 @@ class Transformer(ast.NodeTransformer):
 
     def visit_If(self, node):
         frame = self.fresh("conditional")
-        branches = [
-            self._with(
-                self._operation("Then", [], node), self._scope(node.body, 
node, "then"), node
+        condition_stmt, condition = self._cache(self._expression(node.test), 
node.test, "condition")
+        if self.preserve_return:
+            # Construction helpers retain Python return semantics in ordinary 
branches.
+            then_body = self._scope(node.body, node, "then")
+            else_body = self._scope(node.orelse, node, "else")
+            host = self._located(ast.If(copy.deepcopy(condition), then_body, 
else_body), node)
+            definitions = []
+            then_call, else_call = then_body, else_body
+        else:
+            then_def, then_invoke = self._scope(node.body, node, "then", 
return_bindings=True)
+            else_def, else_invoke = self._scope(node.orelse, node, "else", 
return_bindings=True)
+            definitions = [then_def, else_def]
+            then_call, else_call = [then_invoke], [else_invoke]
+            mapping_name = self.fresh("exports")
+            host = self._located(
+                ast.If(
+                    copy.deepcopy(condition),
+                    [self._assign(mapping_name, 
copy.deepcopy(then_invoke.value), node)],
+                    [self._assign(mapping_name, 
copy.deepcopy(else_invoke.value), node)],
+                ),
+                node,
             )
-        ]
+        branches = [self._with(self._operation("Then", [], node), then_call, 
node)]
         if node.orelse:
-            branches.append(
-                self._with(
-                    self._operation("Else", [], node), 
self._scope(node.orelse, node, "else"), node
-                )
-            )
+            branches.append(self._with(self._operation("Else", [], node), 
else_call, node))
         region = self._with(
-            self._operation("If", [self._expression(node.test)], node),
+            self._operation("If", [copy.deepcopy(condition)], node),
             branches,
             node,
             self._name(frame, node, True),
         )
-        return [region, *self._exports(frame, self._assigned_names(node.body + 
node.orelse), node)]
+        native = [region]
+        candidates = self._assigned_names(node.body + node.orelse)
+        if self.preserve_return:
+            exports = self._exports(frame, candidates, node)
+            native.extend(exports)
+            tail = []
+        else:
+            native.append(
+                self._assign(
+                    mapping_name,
+                    self._call(
+                        self.infrastructure_name, "frame_result", 
[self._name(frame, node)], node
+                    ),
+                    node,
+                )
+            )
+            tail = self._exports(frame, candidates, node, 
self._name(mapping_name, node))
+        dispatch = self._located(
+            ast.If(
+                self._call(
+                    self.infrastructure_name, "is_python_bool", 
[copy.deepcopy(condition)], node
+                ),
+                [host],
+                native,
+            ),
+            node,
+        )
+        return [condition_stmt, *definitions, dispatch, *tail]
 
     def visit_For(self, node):
         if node.orelse:
diff --git a/python/tvm/tirx/script/builder_v2/__init__.py 
b/python/tvm/tirx/script/builder_v2/__init__.py
index 4b4746f0c0..4d9d7cb6aa 100644
--- a/python/tvm/tirx/script/builder_v2/__init__.py
+++ b/python/tvm/tirx/script/builder_v2/__init__.py
@@ -207,7 +207,7 @@ def _check_unterminated():
         last = statements[-1]
         if isinstance(last, _tir.Return | _tir.Break | _tir.Continue) or (
             isinstance(last, _tir.Evaluate)
-            and isinstance(last.value, _tir.Call)
+            and isinstance(last.value, _ir.Call)
             and isinstance(last.value.op, _ir.Op)
             and last.value.op.name in ("tirx.break_loop", "tirx.continue_loop")
         ):
@@ -503,6 +503,9 @@ def logical_and(*values):
     """Construct scalar/vector conjunction, preserving ordinary Python 
values."""
     if not values:
         raise TypeError("logical_and requires at least one operand")
+    values = [
+        value.asobject() if isinstance(value, _ffi.ObjectConvertible) else 
value for value in values
+    ]
     result = values[0]
     for value in values[1:]:
         if not isinstance(result, _ir.Expr) and not isinstance(value, 
_ir.Expr):
@@ -517,6 +520,9 @@ def logical_or(*values):
     """Construct scalar/vector disjunction, preserving ordinary Python 
values."""
     if not values:
         raise TypeError("logical_or requires at least one operand")
+    values = [
+        value.asobject() if isinstance(value, _ffi.ObjectConvertible) else 
value for value in values
+    ]
     result = values[0]
     for value in values[1:]:
         if not isinstance(result, _ir.Expr) and not isinstance(value, 
_ir.Expr):
@@ -529,11 +535,15 @@ def logical_or(*values):
 
 def logical_not(value):
     """Construct IR negation without coercing an IR expression to Python 
bool."""
+    if isinstance(value, _ffi.ObjectConvertible):
+        value = value.asobject()
     return _tir.Not(value) if isinstance(value, _ir.Expr) else not value
 
 
 def select(condition, true_value, false_value):
     """Construct a conditional expression or select an ordinary Python 
value."""
+    if isinstance(condition, _ffi.ObjectConvertible):
+        condition = condition.asobject()
     if not isinstance(condition, _ir.Expr):
         return true_value if condition else false_value
     return _tir.if_then_else(condition, true_value, false_value)

Reply via email to