This is an automated email from the ASF dual-hosted git repository.
spectrometerHBH pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 2647a19cc3 [IR][TIRX] Add first-class tuple expressions (#20168)
2647a19cc3 is described below
commit 2647a19cc39965e39033904f42e934a76d427d53
Author: Tianqi Chen <[email protected]>
AuthorDate: Sun Aug 23 16:41:20 2026 -0400
[IR][TIRX] Add first-class tuple expressions (#20168)
Tuple expressions are independently useful across IR dialects, while
their current Relax ownership prevents TIRX from representing
tuple-valued internal programs directly.
This change:
- promotes Tuple and TupleGetItem to core IR with ir.Tuple and
ir.TupleGetItem runtime keys while preserving Relax source aliases and
legacy FFI constructors;
- converts explicit TIRX list and tuple literals at internal T.let
bindings, without exposing tuple-valued PrimFunc return boundaries or
changing builder-owned host sequences;
- closes TIRX visitor, mutator, deep-equality, path traversal, printer,
and Python functor support over the new values;
- updates existing reflection coverage and adds a minimal internal TIRX
parser/printer round trip.
Validation:
- full Ninja rebuild after the public-header changes
- changed-file pre-commit hooks
- focused existing IR, Relax, and TIRX suites: 255 passed
---
docs/reference/api/python/relax/relax.rst | 2 +-
include/tvm/ir/expr.h | 62 ++++++++++++++++++-
include/tvm/relax/expr.h | 79 ++-----------------------
include/tvm/tirx/expr_functor.h | 8 +++
python/tvm/ir/__init__.py | 12 +++-
python/tvm/ir/expr.py | 57 ++++++++++++++++++
python/tvm/relax/expr.py | 63 +-------------------
python/tvm/tirx/expr_functor.py | 35 ++++++++++-
python/tvm/tirx/script/builder/ir.py | 13 ++++
python/tvm/tirx/script/parser/parser.py | 29 +++++++++
src/ir/expr.cc | 50 ++++++++++++++++
src/relax/ir/expr.cc | 55 -----------------
src/tirx/analysis/deep_equal.cc | 11 ++++
src/tirx/ir/expr_functor.cc | 18 ++++++
src/tirx/ir/py_functor.cc | 4 ++
src/tirx/ir/tir_visitor_with_path.cc | 8 +++
src/tirx/ir/tir_visitor_with_path.h | 6 ++
src/tirx/script/printer/expr.cc | 12 ++++
tests/python/ir/test_node_reflection.py | 2 +-
tests/python/tirx-base/test_tir_constructor.py | 18 ++++++
tests/python/tirx-base/test_tir_expr_functor.py | 43 +++++++++++++-
tests/python/tirx-transform/test_tir_functor.py | 18 ++++++
tests/python/tirx/test_parser_printer.py | 29 +++++++++
23 files changed, 439 insertions(+), 195 deletions(-)
diff --git a/docs/reference/api/python/relax/relax.rst
b/docs/reference/api/python/relax/relax.rst
index fefd074f00..28cdc5e32b 100644
--- a/docs/reference/api/python/relax/relax.rst
+++ b/docs/reference/api/python/relax/relax.rst
@@ -20,4 +20,4 @@ tvm.relax
.. automodule:: tvm.relax
:members:
:imported-members:
- :exclude-members: BlockBuilder, Call, Var, Span, GlobalVar, SourceName,
TupleType, Type, FuncType
+ :exclude-members: BlockBuilder, Call, Tuple, TupleGetItem, Var, Span,
GlobalVar, SourceName, TupleType, Type, FuncType
diff --git a/include/tvm/ir/expr.h b/include/tvm/ir/expr.h
index 4d614586d5..21fa304c3f 100644
--- a/include/tvm/ir/expr.h
+++ b/include/tvm/ir/expr.h
@@ -38,13 +38,73 @@
#include <limits>
#include <optional>
#include <string>
-#include <type_traits>
namespace tvm {
// Forward-declare VirtualDevice to avoid circular imports.
class VirtualDevice;
+/*! \brief Tuple container */
+class TupleNode : public ExprNode {
+ public:
+ /*! \brief The fields of the tuple. */
+ ffi::Array<Expr> fields;
+
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<TupleNode>().def_ro("fields", &TupleNode::fields);
+ }
+
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("ir.Tuple", TupleNode, ExprNode);
+};
+
+/*! \brief Managed reference to TupleNode. */
+class Tuple : public Expr {
+ public:
+ /*!
+ * \brief Construct a tuple from its fields.
+ * \param fields The fields of the tuple.
+ * \param span The source span of the expression.
+ */
+ TVM_DLL explicit Tuple(ffi::Array<Expr> fields, Span span = Span());
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(Tuple, Expr, TupleNode);
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleNode);
+};
+
+/*! \brief Get the index-th field out of a tuple. */
+class TupleGetItemNode : public ExprNode {
+ public:
+ /*! \brief The tuple expression. */
+ Expr tuple;
+ /*! \brief The field index. */
+ int index;
+
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<TupleGetItemNode>()
+ .def_ro("tuple_value", &TupleGetItemNode::tuple)
+ .def_ro("index", &TupleGetItemNode::index);
+ }
+
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("ir.TupleGetItem", TupleGetItemNode,
ExprNode);
+};
+
+/*! \brief Managed reference to TupleGetItemNode. */
+class TupleGetItem : public Expr {
+ public:
+ /*!
+ * \brief Construct a tuple field projection.
+ * \param tuple The tuple to get an element from.
+ * \param index The field index.
+ * \param span The source span of the expression.
+ */
+ TVM_DLL TupleGetItem(Expr tuple, int index, Span span = Span());
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(TupleGetItem, Expr,
TupleGetItemNode);
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleGetItemNode);
+};
+
/*!
* \brief add operator
*
diff --git a/include/tvm/relax/expr.h b/include/tvm/relax/expr.h
index d534c3948e..aadc101560 100644
--- a/include/tvm/relax/expr.h
+++ b/include/tvm/relax/expr.h
@@ -36,80 +36,11 @@
namespace tvm {
namespace relax {
-/*! \brief Tuple container */
-class TupleNode : public ExprNode {
- public:
- /*! \brief the fields of the tuple */
- tvm::ffi::Array<Expr> fields;
-
- static void RegisterReflection() {
- namespace refl = tvm::ffi::reflection;
- refl::ObjectDef<TupleNode>().def_ro("fields", &TupleNode::fields);
- }
- TVM_FFI_DECLARE_OBJECT_INFO_FINAL("relax.expr.Tuple", TupleNode, ExprNode);
-};
-
-class Tuple : public Expr {
- public:
- /*!
- * \brief The constructor
- * \param fields The fields of a tuple.
- * \param span The source span of the expression.
- */
- TVM_DLL explicit Tuple(tvm::ffi::Array<Expr> fields, Span span = Span());
-
- /*!
- * \brief Utility constructor to handle conversion to relax::Expr
- *
- * If the calling scope already has an array of a specific type of
- * relax expression (e.g. `ffi::Array<Var>`), it must be converted
- * into an array of base type. This constructor handles the
- * conversion to the base `ffi::Array<relax::Expr>`.
- *
- * \tparam ExprType The type of relax expression passed in as an argument.
- *
- * \param fields The fields of a tuple.
- *
- * \param span The source span of the expression.
- */
- template <typename ExprType, typename =
std::enable_if_t<std::is_base_of_v<Expr, ExprType>>>
- TVM_DLL explicit Tuple(tvm::ffi::Array<ExprType> fields, Span span = Span())
- : Tuple(fields.Map([](const ExprType& expr) -> Expr { return expr; }),
span) {}
-
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(Tuple, Expr, TupleNode);
- TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleNode);
-};
-
-/*! \brief Get index-th field out of a tuple. */
-class TupleGetItemNode : public ExprNode {
- public:
- /*! \brief The tuple Expression */
- Expr tuple;
- /*! \brief which value to get */
- int index;
-
- static void RegisterReflection() {
- namespace refl = tvm::ffi::reflection;
- refl::ObjectDef<TupleGetItemNode>()
- .def_ro("tuple_value", &TupleGetItemNode::tuple)
- .def_ro("index", &TupleGetItemNode::index);
- }
- TVM_FFI_DECLARE_OBJECT_INFO_FINAL("relax.expr.TupleGetItem",
TupleGetItemNode, ExprNode);
-};
-
-class TupleGetItem : public Expr {
- public:
- /*!
- * \brief The constructor
- * \param tuple The tuple to get an element from.
- * \param index The index for extracting a value in the tuple.
- * \param span The source span of the expression.
- */
- TVM_DLL TupleGetItem(Expr tuple, int index, Span span = Span());
-
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(TupleGetItem, Expr,
TupleGetItemNode);
- TVM_DEFINE_OBJECT_REF_COW_METHOD(TupleGetItemNode);
-};
+// Compatibility aliases. Tuple expressions are defined in the common IR.
+using ::tvm::Tuple;
+using ::tvm::TupleGetItem;
+using ::tvm::TupleGetItemNode;
+using ::tvm::TupleNode;
/*! \brief A shape expression which allows users to construct a shape
containing PrimExpr.
*/
diff --git a/include/tvm/tirx/expr_functor.h b/include/tvm/tirx/expr_functor.h
index d95363a2bc..170c3a499d 100644
--- a/include/tvm/tirx/expr_functor.h
+++ b/include/tvm/tirx/expr_functor.h
@@ -117,6 +117,8 @@ class ExprFunctor<R(const Expr& n, Args...)> {
virtual R VisitExpr_(const VarNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const BufferLoadNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const ProducerLoadNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const TupleNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const TupleGetItemNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const LetNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const CallNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const AddNode* op, Args... args) EXPR_FUNCTOR_DEFAULT;
@@ -159,6 +161,8 @@ class ExprFunctor<R(const Expr& n, Args...)> {
IR_EXPR_FUNCTOR_DISPATCH(VarNode);
IR_EXPR_FUNCTOR_DISPATCH(BufferLoadNode);
IR_EXPR_FUNCTOR_DISPATCH(ProducerLoadNode);
+ IR_EXPR_FUNCTOR_DISPATCH(TupleNode);
+ IR_EXPR_FUNCTOR_DISPATCH(TupleGetItemNode);
IR_EXPR_FUNCTOR_DISPATCH(LetNode);
IR_EXPR_FUNCTOR_DISPATCH(CallNode);
IR_EXPR_FUNCTOR_DISPATCH(AddNode);
@@ -209,6 +213,8 @@ class TVM_DLL ExprVisitor : public ExprFunctor<void(const
Expr&)> {
void VisitExpr_(const VarNode* op) override;
void VisitExpr_(const BufferLoadNode* op) override;
void VisitExpr_(const ProducerLoadNode* op) override;
+ void VisitExpr_(const TupleNode* op) override;
+ void VisitExpr_(const TupleGetItemNode* op) override;
void VisitExpr_(const LetNode* op) override;
void VisitExpr_(const CallNode* op) override;
void VisitExpr_(const AddNode* op) override;
@@ -255,6 +261,8 @@ class TVM_DLL ExprMutator : protected
ExprFunctor<Expr(const Expr&)> {
Expr VisitExpr_(const VarNode* op) override;
Expr VisitExpr_(const BufferLoadNode* op) override;
Expr VisitExpr_(const ProducerLoadNode* op) override;
+ Expr VisitExpr_(const TupleNode* op) override;
+ Expr VisitExpr_(const TupleGetItemNode* op) override;
Expr VisitExpr_(const LetNode* op) override;
Expr VisitExpr_(const CallNode* op) override;
Expr VisitExpr_(const AddNode* op) override;
diff --git a/python/tvm/ir/__init__.py b/python/tvm/ir/__init__.py
index 3dbb9b234f..d907119372 100644
--- a/python/tvm/ir/__init__.py
+++ b/python/tvm/ir/__init__.py
@@ -34,7 +34,17 @@ from .base import (
# Register Type before Expr. Expr's reflected ``ty`` field otherwise creates
# an auto-generated Type wrapper before the concrete Python class is available.
from .type import FuncType, PointerType, PrimType, TupleType, Type
-from .expr import Call, Expr, GlobalVar, Range, Var, is_prim_expr, is_prim_var
+from .expr import (
+ Call,
+ Expr,
+ GlobalVar,
+ Range,
+ Tuple,
+ TupleGetItem,
+ Var,
+ is_prim_expr,
+ is_prim_var,
+)
from .function import BaseFunc, CallingConv
from .global_info import GlobalInfo, DummyGlobalInfo, VDevice
from .module import IRModule
diff --git a/python/tvm/ir/expr.py b/python/tvm/ir/expr.py
index 4c28b7e7f6..576a02223f 100644
--- a/python/tvm/ir/expr.py
+++ b/python/tvm/ir/expr.py
@@ -328,6 +328,63 @@ class _ExprWithOp(Expr, Scriptable):
return result
+@tvm_ffi.register_object("ir.Tuple")
+class Tuple(_ExprWithOp):
+ """Tuple expression that groups several fields together.
+
+ Parameters
+ ----------
+ fields : list[Expr] | tuple[Expr, ...]
+ The fields in the tuple.
+
+ span : Span | None
+ Span that points to the original source code.
+ """
+
+ fields: list[Expr]
+ span: Span | None
+
+ def __init__(self, fields: list[Expr] | tuple[Expr, ...], span: Span |
None = None):
+ if isinstance(fields, Tuple):
+ fields = fields.fields
+ elif isinstance(getattr(fields, "ty", None), tvm.ir.TupleType):
+ fields = [*fields]
+
+ self.__init_handle_by_constructor__(_ffi_api.Tuple, fields, span)
+
+ def __getitem__(self, index: int) -> Expr:
+ if index >= len(self) or index < -len(self):
+ raise IndexError("Tuple index out of range")
+ return self.fields[index]
+
+ def __len__(self) -> int:
+ return len(self.fields)
+
+
+@tvm_ffi.register_object("ir.TupleGetItem")
+class TupleGetItem(_ExprWithOp):
+ """Get the index-th item from a tuple.
+
+ Parameters
+ ----------
+ tuple_value : Expr
+ The input tuple expression.
+
+ index : int
+ The field index.
+
+ span : Span | None
+ Span that points to the original source code.
+ """
+
+ tuple_value: Expr
+ index: int
+ span: Span | None
+
+ def __init__(self, tuple_value: Expr, index: int, span: Span | None =
None):
+ self.__init_handle_by_constructor__(_ffi_api.TupleGetItem,
tuple_value, index, span)
+
+
@tvm_ffi.register_object("ir.Call")
class Call(_ExprWithOp):
"""Core function call node."""
diff --git a/python/tvm/relax/expr.py b/python/tvm/relax/expr.py
index 8c7cdcf992..3c983da029 100644
--- a/python/tvm/relax/expr.py
+++ b/python/tvm/relax/expr.py
@@ -270,66 +270,9 @@ class If(ExprWithOp):
)
-@tvm_ffi.register_object("relax.expr.Tuple")
-class Tuple(ExprWithOp):
- """Tuple expression that groups several fields together.
-
- Parameters
- ----------
- fields : Union[List[Expr], typing.Tuple[Expr, ...]]
- The fields in the tuple.
-
- span: Optional[Span]
- Span that points to original source code
- """
-
- fields: list[Expr]
- span: Span | None
-
- def __init__(self, fields: list[Expr] | tuple[Expr, ...], span: Span |
None = None):
- if isinstance(fields, tvm.relax.Tuple):
- fields = fields.fields
- elif isinstance(getattr(fields, "ty", None), tvm.relax.TupleType):
- fields = [*fields]
-
- self.__init_handle_by_constructor__(_ffi_api.Tuple, fields, span) #
type: ignore
-
- def __getitem__(self, index: int) -> Expr:
- if index >= len(self) or index < -len(self):
- raise IndexError("Tuple index out of range")
- return self.fields[index]
-
- def __len__(self) -> int:
- return len(self.fields)
-
-
-@tvm_ffi.register_object("relax.expr.TupleGetItem")
-class TupleGetItem(ExprWithOp):
- """Get index-th item from a tuple.
-
- Parameters
- ----------
- tuple_value: Expr
- The input tuple expression.
-
- index: int
- The index.
-
- span: Optional[Span]
- Span that points to original source code
- """
-
- tuple_value: Expr
- index: int
- span: Span | None
-
- def __init__(self, tuple_value: Expr, index: int, span: Span | None =
None):
- self.__init_handle_by_constructor__(
- _ffi_api.TupleGetItem,
- tuple_value,
- index,
- span, # type: ignore
- )
+# Compatibility aliases. Tuple expressions are owned by the common IR.
+Tuple = tvm.ir.Tuple
+TupleGetItem = tvm.ir.TupleGetItem
@tvm_ffi.register_object("relax.expr.ShapeExpr")
diff --git a/python/tvm/tirx/expr_functor.py b/python/tvm/tirx/expr_functor.py
index 819a630d41..27dc87c50a 100644
--- a/python/tvm/tirx/expr_functor.py
+++ b/python/tvm/tirx/expr_functor.py
@@ -24,7 +24,7 @@ from collections.abc import Callable
from typing import TypeVar
import tvm
-from tvm.ir import Expr, Range
+from tvm.ir import Expr, Range, Tuple, TupleGetItem
from tvm.tirx import IterVar
T = TypeVar("T")
@@ -52,6 +52,8 @@ class ExprFunctor:
"tirx.Var": self.visit_var_,
"tirx.BufferLoad": self.visit_buffer_load_,
"tirx.ProducerLoad": self.visit_producer_load_,
+ "tirx.Tuple": self.visit_tuple_,
+ "tirx.TupleGetItem": self.visit_tuple_get_item_,
"tirx.Let": self.visit_let_,
"tirx.Call": self.visit_call_,
"tirx.Add": self.visit_add_,
@@ -121,6 +123,14 @@ class ExprFunctor:
"""Default visitor for ProducerLoad node."""
return self.visit_expr_default_(op)
+ def visit_tuple_(self, op):
+ """Default visitor for Tuple node."""
+ return self.visit_expr_default_(op)
+
+ def visit_tuple_get_item_(self, op):
+ """Default visitor for TupleGetItem node."""
+ return self.visit_expr_default_(op)
+
def visit_let_(self, op):
"""Default visitor for Let node."""
return self.visit_expr_default_(op)
@@ -284,6 +294,14 @@ class ExprVisitor(ExprFunctor):
_visit_array(op.indices, _visit_indices)
+ def visit_tuple_(self, op):
+ """Visitor implementation for Tuple."""
+ _visit_array(op.fields, self.visit_expr)
+
+ def visit_tuple_get_item_(self, op):
+ """Visitor implementation for TupleGetItem."""
+ self.visit_expr(op.tuple_value)
+
def visit_let_(self, op):
"""Visitor implementation for Let."""
self.visit_expr(op.value)
@@ -464,6 +482,21 @@ class ExprMutator(ExprFunctor):
else:
return tvm.tirx.ProducerLoad(op.producer, indices)
+ def visit_tuple_(self, op):
+ """Mutator implementation for Tuple."""
+ fields = [self.visit_expr(field) for field in op.fields]
+
+ if all(old_field is new_field for old_field, new_field in
zip(op.fields, fields)):
+ return op
+ return Tuple(fields, op.span)
+
+ def visit_tuple_get_item_(self, op):
+ """Mutator implementation for TupleGetItem."""
+ tuple_value = self.visit_expr(op.tuple_value)
+ if tuple_value is op.tuple_value:
+ return op
+ return TupleGetItem(tuple_value, op.index, op.span)
+
def visit_let_(self, op):
"""Mutator implementation for Let."""
var = self.visit_var_(op.var)
diff --git a/python/tvm/tirx/script/builder/ir.py
b/python/tvm/tirx/script/builder/ir.py
index bf2005befa..e1b87e1f25 100644
--- a/python/tvm/tirx/script/builder/ir.py
+++ b/python/tvm/tirx/script/builder/ir.py
@@ -449,6 +449,18 @@ def func_ret(ret_type: Type | None) -> Type:
return _ffi_api.FuncRet(ret_type) # type: ignore[attr-defined] # pylint:
disable=no-member
+def Tuple(*fields: Type) -> Type: # pylint: disable=invalid-name
+ """Construct a tuple type for a TIRx function or binding annotation."""
+ normalized_fields = []
+ for field in fields:
+ if callable(field) and not isinstance(field, Expr):
+ field = field()
+ if isinstance(field, Expr):
+ field = field.ty
+ normalized_fields.append(field)
+ return ir.TupleType(normalized_fields)
+
+
def match_buffer(
param: Var | BufferLoad | BufferRegion,
shape: list[Expr] | tuple[Expr] | Expr | Integral = None,
@@ -3355,6 +3367,7 @@ __all__ = [
"func_name",
"func_attr",
"func_ret",
+ "Tuple",
"match_buffer",
"sblock",
"block_name_suffix_context",
diff --git a/python/tvm/tirx/script/parser/parser.py
b/python/tvm/tirx/script/parser/parser.py
index ca881ac1c7..2efef79d07 100644
--- a/python/tvm/tirx/script/parser/parser.py
+++ b/python/tvm/tirx/script/parser/parser.py
@@ -163,6 +163,34 @@ def bind_for_value(self: Parser, node: doc.expr, var_name:
str, value: Any) -> A
raise NotImplementedError
+def _convert_tuple_literal(self: Parser, node: doc.expr, value: Any) -> Any:
+ """Convert an explicit list/tuple T.let value to a core IR Tuple.
+
+ The generic evaluator deliberately keeps Python containers intact because
+ many TIRx builder APIs use them structurally. Conversion is restricted to
+ the internal immutable-binding boundary that explicitly opts in below.
+ """
+ if not isinstance(node, doc.List | doc.Tuple):
+ return value
+
+ def convert_field(field: Any) -> Expr:
+ if isinstance(field, list | tuple):
+ return tvm.ir.Tuple([convert_field(item) for item in field])
+ if isinstance(field, Expr):
+ return field
+ if isinstance(field, str):
+ return tvm.tirx.StringImm(field)
+ if isinstance(field, bool | int | float):
+ return tvm.tirx.const(field)
+ self.report_error(
+ node,
+ f"Tuple fields must be expressions or scalar literals, got
{type(field).__name__}",
+ )
+ raise NotImplementedError
+
+ return convert_field(value)
+
+
def bind_assign_value(
self: Parser,
node: doc.expr,
@@ -759,6 +787,7 @@ def visit_ann_assign(self: Parser, node: doc.AnnAssign) ->
None:
# T.let or T.let[type] -> immutable Bind var
if rhs is None:
self.report_error(node, "T.let annotation requires a value")
+ rhs = _convert_tuple_literal(self, node.value, rhs)
if not isinstance(rhs, Expr):
if isinstance(rhs, str):
rhs = tvm.tirx.StringImm(rhs)
diff --git a/src/ir/expr.cc b/src/ir/expr.cc
index d44cdb06f6..10714b7953 100644
--- a/src/ir/expr.cc
+++ b/src/ir/expr.cc
@@ -42,11 +42,61 @@ TVM_FFI_STATIC_INIT_BLOCK() {
VarNode::RegisterReflection();
GlobalVarNode::RegisterReflection();
CallNode::RegisterReflection();
+ TupleNode::RegisterReflection();
+ TupleGetItemNode::RegisterReflection();
IntImmNode::RegisterReflection();
FloatImmNode::RegisterReflection();
RangeNode::RegisterReflection();
}
+Tuple::Tuple(ffi::Array<Expr> fields, Span span) {
+ ffi::Optional<Type> tuple_ty = [&]() -> ffi::Optional<Type> {
+ ffi::Array<Type> field_ty;
+ for (const Expr& field : fields) {
+ if (field->ty.IsMissing()) {
+ return std::nullopt;
+ }
+ field_ty.push_back(field->ty);
+ }
+ return TupleType(field_ty);
+ }();
+
+ ffi::ObjectPtr<TupleNode> node = ffi::make_object<TupleNode>();
+ node->fields = std::move(fields);
+ node->span = std::move(span);
+ if (tuple_ty.has_value()) {
+ node->ty = tuple_ty.value();
+ }
+ data_ = std::move(node);
+}
+
+TupleGetItem::TupleGetItem(Expr tuple, int index, Span span) {
+ TVM_FFI_ICHECK_GE(index, 0) << "Index out of bounds: Tuple " << tuple
+ << " cannot be accessed with negative index " <<
index;
+ ffi::ObjectPtr<TupleGetItemNode> node = ffi::make_object<TupleGetItemNode>();
+ if (const auto* tuple_type = tuple->ty.as<TupleTypeNode>()) {
+ TVM_FFI_ICHECK_LT(index, tuple_type->fields.size())
+ << "Index out of bounds: Tuple " << tuple << " is of size " <<
tuple_type->fields.size()
+ << ", and cannot be accessed with index " << index;
+ node->ty = tuple_type->fields[index];
+ }
+ node->tuple = std::move(tuple);
+ node->index = index;
+ node->span = std::move(span);
+ data_ = std::move(node);
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ refl::GlobalDef()
+ .def("ir.Tuple", [](ffi::Array<Expr> fields, Span span) { return
Tuple(fields, span); })
+ .def("ir.TupleGetItem",
+ [](Expr tuple, int index, Span span) { return TupleGetItem(tuple,
index, span); })
+ .def("relax.Tuple", [](ffi::Array<Expr> fields, Span span) { return
Tuple(fields, span); })
+ .def("relax.TupleGetItem",
+ [](Expr tuple, int index, Span span) { return TupleGetItem(tuple,
index, span); });
+}
+
PrimExpr::PrimExpr(Call call) :
PrimExpr(std::move(call).as_or_throw<PrimExpr>()) {}
PrimExpr::PrimExpr(int32_t value) : PrimExpr(IntImm::Int32(value)) {}
diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc
index 3242fb2224..b21cdbad87 100644
--- a/src/relax/ir/expr.cc
+++ b/src/relax/ir/expr.cc
@@ -28,8 +28,6 @@ namespace tvm {
namespace relax {
TVM_FFI_STATIC_INIT_BLOCK() {
- TupleNode::RegisterReflection();
- TupleGetItemNode::RegisterReflection();
ShapeExprNode::RegisterReflection();
BindingNode::RegisterReflection();
DataflowVarNode::RegisterReflection();
@@ -62,59 +60,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
});
}
-Tuple::Tuple(tvm::ffi::Array<Expr> fields, Span span) {
- ffi::Optional<Type> tuple_ty = [&]() -> ffi::Optional<Type> {
- ffi::Array<Type> field_ty;
- for (const auto& field : fields) {
- if (!field->ty.IsMissing()) {
- field_ty.push_back(GetType(field));
- } else {
- return std::nullopt;
- }
- }
- return TupleType(field_ty);
- }();
-
- ffi::ObjectPtr<TupleNode> n = ffi::make_object<TupleNode>();
- n->fields = std::move(fields);
- n->span = std::move(span);
- if (tuple_ty.has_value()) {
- n->ty = tuple_ty.value();
- }
- data_ = std::move(n);
-}
-
-TVM_FFI_STATIC_INIT_BLOCK() {
- namespace refl = tvm::ffi::reflection;
- refl::GlobalDef().def(
- "relax.Tuple", [](tvm::ffi::Array<Expr> fields, Span span) { return
Tuple(fields, span); });
-}
-
-TupleGetItem::TupleGetItem(Expr tuple, int index, Span span) {
- TVM_FFI_ICHECK_GE(index, 0) << "Index out of bounds: Tuple " << tuple
- << " cannot be accessed with negative index " <<
index;
- ffi::ObjectPtr<TupleGetItemNode> n = ffi::make_object<TupleGetItemNode>();
-
- if (auto* tuple_info = tuple->ty.as<TupleTypeNode>()) {
- TVM_FFI_ICHECK_LT(index, tuple_info->fields.size())
- << "Index out of bounds: Tuple " << tuple << " is of size " <<
tuple_info->fields.size()
- << ", and cannot be accessed with index " << index;
- auto ty = tuple_info->fields[index];
- n->ty = ty;
- }
- n->tuple = std::move(tuple);
- n->index = index;
- n->span = std::move(span);
- data_ = std::move(n);
-}
-
-TVM_FFI_STATIC_INIT_BLOCK() {
- namespace refl = tvm::ffi::reflection;
- refl::GlobalDef().def("relax.TupleGetItem", [](Expr tuple, int index, Span
span) {
- return TupleGetItem(tuple, index, span);
- });
-}
-
ShapeExpr::ShapeExpr(ffi::Array<PrimExpr> values, Span span) {
ffi::ObjectPtr<ShapeExprNode> n = ffi::make_object<ShapeExprNode>();
diff --git a/src/tirx/analysis/deep_equal.cc b/src/tirx/analysis/deep_equal.cc
index d40baa2877..48fcb120a2 100644
--- a/src/tirx/analysis/deep_equal.cc
+++ b/src/tirx/analysis/deep_equal.cc
@@ -80,6 +80,17 @@ class ExprDeepEqualChecker : private ExprFunctor<bool(const
Expr&, const PrimExp
if (lhs.as<VarNode>()) {
return false;
}
+ if (auto* lhs_tuple = lhs.as<TupleNode>()) {
+ auto* rhs_tuple = rhs.as<TupleNode>();
+ return ffi::StructuralEqual()(lhs_tuple->ty, rhs_tuple->ty) &&
+ ArrayDeepEqual(lhs_tuple->fields, rhs_tuple->fields);
+ }
+ if (auto* lhs_get_item = lhs.as<TupleGetItemNode>()) {
+ auto* rhs_get_item = rhs.as<TupleGetItemNode>();
+ return lhs_get_item->index == rhs_get_item->index &&
+ ffi::StructuralEqual()(lhs_get_item->ty, rhs_get_item->ty) &&
+ VisitExpr(lhs_get_item->tuple, rhs_get_item->tuple);
+ }
if (auto* lhs_call = lhs.as<CallNode>()) {
auto* rhs_call = rhs.as<CallNode>();
return ffi::StructuralEqual()(lhs_call->ty, rhs_call->ty) &&
diff --git a/src/tirx/ir/expr_functor.cc b/src/tirx/ir/expr_functor.cc
index 2fbd99f62e..e465ac970f 100644
--- a/src/tirx/ir/expr_functor.cc
+++ b/src/tirx/ir/expr_functor.cc
@@ -38,6 +38,12 @@ void ExprVisitor::VisitExpr_(const ProducerLoadNode* op) {
VisitArray(op->indices, [this](const PrimExpr& e) { this->VisitExpr(e); });
}
+void ExprVisitor::VisitExpr_(const TupleNode* op) {
+ VisitArray(op->fields, [this](const Expr& e) { this->VisitExpr(e); });
+}
+
+void ExprVisitor::VisitExpr_(const TupleGetItemNode* op) {
this->VisitExpr(op->tuple); }
+
void ExprVisitor::VisitExpr_(const LetNode* op) {
this->VisitExpr(op->value);
this->VisitExpr(op->body);
@@ -131,6 +137,18 @@ Expr ExprMutator::VisitExpr_(const ProducerLoadNode* op) {
}
}
+Expr ExprMutator::VisitExpr_(const TupleNode* op) {
+ ffi::Array<Expr> fields =
+ op->fields.Map([this](const Expr& field) { return
this->VisitExpr(field); });
+ return fields.same_as(op->fields) ? ffi::GetRef<tvm::Tuple>(op) :
tvm::Tuple(fields, op->span);
+}
+
+Expr ExprMutator::VisitExpr_(const TupleGetItemNode* op) {
+ Expr tuple_value = this->VisitExpr(op->tuple);
+ return tuple_value.same_as(op->tuple) ? ffi::GetRef<TupleGetItem>(op)
+ : TupleGetItem(std::move(tuple_value),
op->index, op->span);
+}
+
Expr ExprMutator::VisitExpr_(const LetNode* op) {
PrimExpr value = this->VisitPrimExpr(op->value);
PrimExpr body = this->VisitPrimExpr(op->body);
diff --git a/src/tirx/ir/py_functor.cc b/src/tirx/ir/py_functor.cc
index 3c451e5277..8c20471f7b 100644
--- a/src/tirx/ir/py_functor.cc
+++ b/src/tirx/ir/py_functor.cc
@@ -272,6 +272,8 @@ class PyStmtExprVisitorNode : public ffi::Object, public
StmtExprVisitor {
IR_EXPR_VISITOR_DEFAULT_DISPATCH(VarNode);
IR_EXPR_VISITOR_DEFAULT_DISPATCH(BufferLoadNode);
IR_EXPR_VISITOR_DEFAULT_DISPATCH(ProducerLoadNode);
+ IR_EXPR_VISITOR_DEFAULT_DISPATCH(TupleNode);
+ IR_EXPR_VISITOR_DEFAULT_DISPATCH(TupleGetItemNode);
IR_EXPR_VISITOR_DEFAULT_DISPATCH(LetNode);
IR_EXPR_VISITOR_DEFAULT_DISPATCH(CallNode);
IR_EXPR_VISITOR_DEFAULT_DISPATCH(AddNode);
@@ -621,6 +623,8 @@ class PyStmtExprMutatorNode : public ffi::Object, public
StmtExprMutator {
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(VarNode);
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(BufferLoadNode);
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(ProducerLoadNode);
+ PY_EXPR_MUTATOR_DEFAULT_DISPATCH(TupleNode);
+ PY_EXPR_MUTATOR_DEFAULT_DISPATCH(TupleGetItemNode);
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(LetNode);
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(CallNode);
PY_EXPR_MUTATOR_DEFAULT_DISPATCH(AddNode);
diff --git a/src/tirx/ir/tir_visitor_with_path.cc
b/src/tirx/ir/tir_visitor_with_path.cc
index 4b39167651..882134f600 100644
--- a/src/tirx/ir/tir_visitor_with_path.cc
+++ b/src/tirx/ir/tir_visitor_with_path.cc
@@ -361,6 +361,14 @@ void TIRVisitorWithPath::VisitExpr_(const
ProducerLoadNode* op, AccessPath path)
Visit(op->indices, path->Attr("indices"));
}
+void TIRVisitorWithPath::VisitExpr_(const TupleNode* op, AccessPath path) {
+ Visit(op->fields, path->Attr("fields"));
+}
+
+void TIRVisitorWithPath::VisitExpr_(const TupleGetItemNode* op, AccessPath
path) {
+ Visit(op->tuple, path->Attr("tuple"));
+}
+
void TIRVisitorWithPath::VisitExpr_(const LetNode* op, AccessPath path) {
Visit(op->value, path->Attr("value"));
auto context = WithDef(op->var, path->Attr("var"));
diff --git a/src/tirx/ir/tir_visitor_with_path.h
b/src/tirx/ir/tir_visitor_with_path.h
index e9735a89b8..d80d44c0cb 100644
--- a/src/tirx/ir/tir_visitor_with_path.h
+++ b/src/tirx/ir/tir_visitor_with_path.h
@@ -62,6 +62,10 @@ class TIRVisitorWithPath : protected ExprFunctor<void(const
Expr&, ffi::reflecti
VisitExpr_(var, path);
} else if (auto* call = obj.as<CallNode>()) {
VisitExpr_(call, path);
+ } else if (auto* tuple = obj.as<TupleNode>()) {
+ VisitExpr_(tuple, path);
+ } else if (auto* tuple_get_item = obj.as<TupleGetItemNode>()) {
+ VisitExpr_(tuple_get_item, path);
} else {
TVM_FFI_THROW(TypeError) << "Unsupported non-primitive TIR expression "
<< obj.GetTypeKey();
}
@@ -147,6 +151,8 @@ class TIRVisitorWithPath : protected ExprFunctor<void(const
Expr&, ffi::reflecti
void VisitExpr_(const VarNode* op, ffi::reflection::AccessPath path)
override;
void VisitExpr_(const BufferLoadNode* op, ffi::reflection::AccessPath path)
override;
void VisitExpr_(const ProducerLoadNode* op, ffi::reflection::AccessPath
path) override;
+ void VisitExpr_(const TupleNode* op, ffi::reflection::AccessPath path)
override;
+ void VisitExpr_(const TupleGetItemNode* op, ffi::reflection::AccessPath
path) override;
void VisitExpr_(const LetNode* op, ffi::reflection::AccessPath path)
override;
void VisitExpr_(const CallNode* op, ffi::reflection::AccessPath path)
override;
void VisitExpr_(const AddNode* op, ffi::reflection::AccessPath path)
override;
diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc
index f0a46503e4..b7b4fc6615 100644
--- a/src/tirx/script/printer/expr.cc
+++ b/src/tirx/script/printer/expr.cc
@@ -24,6 +24,18 @@ namespace tvm {
namespace script {
namespace printer {
+TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable)
+ .set_dispatch<Tuple>("tirx", [](Tuple tuple, AccessPath tuple_p,
IRDocsifier d) -> Doc {
+ return TupleDoc(d->AsDoc<ListDoc>(tuple->fields,
tuple_p->Attr("fields"))->elements);
+ });
+
+TVM_STATIC_IR_FUNCTOR(IRDocsifier, vtable)
+ .set_dispatch<TupleGetItem>(
+ "tirx", [](TupleGetItem get_item, AccessPath get_item_p, IRDocsifier
d) -> Doc {
+ ExprDoc index = LiteralDoc::Int(get_item->index,
get_item_p->Attr("index"));
+ return d->AsDoc<ExprDoc>(get_item->tuple,
get_item_p->Attr("tuple"))[{index}];
+ });
+
ExprDoc PrintVarCreation(const tirx::Var& var, const AccessPath& var_p, const
IRDocsifier& d) {
Type type = var->ty;
AccessPath type_p = var_p->Attr("ty");
diff --git a/tests/python/ir/test_node_reflection.py
b/tests/python/ir/test_node_reflection.py
index 29e222c6db..b1f0744fb1 100644
--- a/tests/python/ir/test_node_reflection.py
+++ b/tests/python/ir/test_node_reflection.py
@@ -59,7 +59,7 @@ _LEGACY_RELAX_VAR_JSON = """{
{"type": "ffi.String", "data": "legacy"},
{"type": "relax.expr.Var", "data": {"span": 1, "ty": 3, "name_hint": 6}},
{"type": "ffi.Array", "data": [7, 7]},
- {"type": "relax.expr.Tuple", "data": {"span": 1, "ty": 5, "fields": 8}}
+ {"type": "ir.Tuple", "data": {"span": 1, "ty": 5, "fields": 8}}
],
"metadata": {"tvm_version": "0.26.dev0"}
}"""
diff --git a/tests/python/tirx-base/test_tir_constructor.py
b/tests/python/tirx-base/test_tir_constructor.py
index 4be31c0174..6dd7410750 100644
--- a/tests/python/tirx-base/test_tir_constructor.py
+++ b/tests/python/tirx-base/test_tir_constructor.py
@@ -171,6 +171,24 @@ def test_expr_constructor():
)
assert not expr_deep_equal(x_with_attrs, x_with_other_attrs)
+ tuple_arg = tvm.ir.Tuple([attr_arg, tvm.tirx.IntImm("int32", 1)])
+ same_tuple_arg = tvm.ir.Tuple([attr_arg, tvm.tirx.IntImm("int32", 1)])
+ different_tuple_arg = tvm.ir.Tuple([attr_arg, tvm.tirx.IntImm("int32", 2)])
+
+ def call_with(arg):
+ return tvm.ir.Call(
+ "tirx.call_extern",
+ [tvm.tirx.StringImm("tuple_arg"), arg],
+ ret_ty="int32",
+ )
+
+ assert expr_deep_equal(call_with(tuple_arg), call_with(same_tuple_arg))
+ assert not expr_deep_equal(call_with(tuple_arg),
call_with(different_tuple_arg))
+ assert not expr_deep_equal(
+ call_with(tuple_arg),
+ call_with(tvm.ir.TupleGetItem(same_tuple_arg, 0)),
+ )
+
cond0 = tvm.tirx.Var("cond0", "bool")
cond1 = tvm.tirx.Var("cond1", "bool")
inner_if = tvm.ir.Call(
diff --git a/tests/python/tirx-base/test_tir_expr_functor.py
b/tests/python/tirx-base/test_tir_expr_functor.py
index 96bb500111..b342782d48 100644
--- a/tests/python/tirx-base/test_tir_expr_functor.py
+++ b/tests/python/tirx-base/test_tir_expr_functor.py
@@ -18,7 +18,7 @@
import tvm
import tvm.testing
from tvm import tirx as tir
-from tvm.ir import Call, Op
+from tvm.ir import Call, Op, Tuple, TupleGetItem
from tvm.ir.base import assert_structural_equal
from tvm.tirx.expr import (
EQ,
@@ -111,6 +111,19 @@ class ASTPrinter(ExprVisitor):
self.visit_expr(idx)
self.log.pop_scope()
+ def visit_tuple_(self, op: Tuple) -> None:
+ self.log.add("Tuple")
+ self.log.push_scope()
+ for field in op.fields:
+ self.visit_expr(field)
+ self.log.pop_scope()
+
+ def visit_tuple_get_item_(self, op: TupleGetItem) -> None:
+ self.log.add("TupleGetItem")
+ self.log.push_scope()
+ self.visit_expr(op.tuple_value)
+ self.log.pop_scope()
+
def visit_let_(self, op: Let) -> None:
self.log.add("Let")
self.log.push_scope()
@@ -339,6 +352,16 @@ class ASTPostPrinterMutator(ExprMutator):
self.log.add("ProducerLoad")
return result
+ def visit_tuple_(self, op: Tuple) -> tir.Expr:
+ result = super().visit_tuple_(op)
+ self.log.add("Tuple")
+ return result
+
+ def visit_tuple_get_item_(self, op: TupleGetItem) -> tir.Expr:
+ result = super().visit_tuple_get_item_(op)
+ self.log.add("TupleGetItem")
+ return result
+
def visit_let_(self, op: Let) -> tir.Expr:
result = super().visit_let_(op)
self.log.add("Let")
@@ -524,6 +547,24 @@ def test_string_imm():
basic_check(tir.StringImm("hello"), "StringImm", "StringImm")
+def test_tuple():
+ tuple_node = Tuple([n, tir.IntImm("int32", 10)])
+ basic_check(
+ tuple_node,
+ "\n".join(["Tuple", "\tVar", "\tIntImm"]),
+ "\n".join(["Var", "IntImm", "Tuple"]),
+ )
+
+
+def test_tuple_get_item():
+ tuple_get_item = TupleGetItem(Tuple([n, m]), 1)
+ basic_check(
+ tuple_get_item,
+ "\n".join(["TupleGetItem", "\tTuple", "\t\tVar", "\t\tVar"]),
+ "\n".join(["Var", "Var", "Tuple", "TupleGetItem"]),
+ )
+
+
def test_add():
add_node = tir.Add(n, m)
basic_check(add_node, "\n".join(["Add", "\tVar", "\tVar"]),
"\n".join(["Var", "Var", "Add"]))
diff --git a/tests/python/tirx-transform/test_tir_functor.py
b/tests/python/tirx-transform/test_tir_functor.py
index 8b463c19a8..14d439219f 100644
--- a/tests/python/tirx-transform/test_tir_functor.py
+++ b/tests/python/tirx-transform/test_tir_functor.py
@@ -19,6 +19,7 @@
import tvm
import tvm.testing
from tvm import tirx
+from tvm.ir import Tuple, TupleGetItem
from tvm.tirx import (
EQ,
LT,
@@ -415,6 +416,23 @@ def test_nested_expressions():
assert counter.mul_count == 1 # one mul
+def test_tuple_default_traversal_and_mutation():
+ x = Var("x", ty="int32")
+ y = Var("y", ty="int32")
+ expr = TupleGetItem(Tuple([x, y]), 1)
+
+ counter = SimpleExprCounter()
+ counter.visit_expr(expr)
+ assert counter.var_count == 2
+
+ replacer = VariableReplacer({"x": 1, "y": 2})
+ result = replacer.visit_expr(expr)
+ assert isinstance(result, TupleGetItem)
+ assert isinstance(result.tuple_value, Tuple)
+ assert result.tuple_value.fields[0].value == 1
+ assert result.tuple_value.fields[1].value == 2
+
+
def test_simple_mutations():
"""Test simple expression mutations"""
x = Var("x", ty="int32")
diff --git a/tests/python/tirx/test_parser_printer.py
b/tests/python/tirx/test_parser_printer.py
index 0e523c29cd..580d77ce34 100644
--- a/tests/python/tirx/test_parser_printer.py
+++ b/tests/python/tirx/test_parser_printer.py
@@ -1240,6 +1240,35 @@ def test_let_annotation_syntax():
assert_structural_equal(test, from_source(code))
+def test_tuple_let_binding_and_traversal():
+ @T.prim_func
+ def from_list(x: T.int32, y: T.float32) -> T.int32:
+ pair: T.let = [x, (y,)]
+ return pair[0]
+
+ @T.prim_func
+ def from_tuple(x: T.int32, y: T.float32) -> T.int32:
+ pair: T.let = (x, (y,))
+ return pair[0]
+
+ def tuple_value(func):
+ visited = []
+ tvm.tirx.stmt_functor.post_order_visit(func.body, visited.append)
+ bind = next(node for node in visited if isinstance(node,
tvm.tirx.Bind))
+ return bind.value
+
+ list_value = tuple_value(from_list)
+ tuple_value = tuple_value(from_tuple)
+ assert isinstance(list_value, tvm.ir.Tuple)
+ assert isinstance(list_value.fields[1], tvm.ir.Tuple)
+ assert_structural_equal(list_value, tuple_value, map_free_vars=True)
+
+ code = from_list.script()
+ assert "pair: T.let[T.Tuple(T.int32, T.Tuple(T.float32))] = x, (y,)" in
code
+ assert from_source(code).script() == code
+ assert_structural_equal(from_list, from_source(code))
+
+
def test_annotation_syntax_comprehensive():
"""Comprehensive test for scalar annotation, T.let, banned annotations,
and bare assignment."""