This is an automated email from the ASF dual-hosted git repository.
tqchen 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 ac8d5ac45d [REFACTOR][TIRx] Group buffer APIs and remove generic
composition (#20402)
ac8d5ac45d is described below
commit ac8d5ac45d03e80792d29067a265b69ea3f60c65
Author: Tianqi Chen <[email protected]>
AuthorDate: Mon Sep 21 20:50:12 2026 -0400
[REFACTOR][TIRx] Group buffer APIs and remove generic composition (#20402)
Group TIRx buffer variables, loads, and region constructors in
`tirx/expr.h`, and consolidate buffer-related types and their
registrations in `tirx/type.h` and its implementation. Preserve the
shared core `TensorRegion` representation and existing buffer semantics.
Remove the unfinished generic `compose_op` builder, printer,
registration, and dispatch surface. Keep the implemented `binary_chain`,
`binary_reduce`, `unary_reduce`, and `reduce_negate` operators and their
shared helpers.
---
docs/tirx/api/tile.rst | 4 -
docs/tirx/api/tirx.rst | 5 +
docs/tirx/tile_primitives.rst | 3 +-
include/tvm/te/operation.h | 2 +-
include/tvm/tirx/buffer_region.h | 58 -----
include/tvm/tirx/{buffer.h => expr.h} | 149 ++-----------
include/tvm/tirx/function.h | 2 +-
include/tvm/tirx/script/builder/frame.h | 31 ---
include/tvm/tirx/script/builder/ir.h | 11 -
include/tvm/tirx/stmt.h | 3 +-
include/tvm/tirx/tile_primitive.h | 2 -
include/tvm/tirx/type.h | 150 +++++++++++++
.../trn/tile_primitive/compose_op/__init__.py | 1 -
.../trn/tile_primitive/compose_op/compose_op.py | 47 ----
python/tvm/tirx/operator/tile_primitive/ops.py | 25 ---
python/tvm/tirx/script/builder/frame.py | 4 -
python/tvm/tirx/script/builder/tirx.py | 24 +-
python/tvm/tirx/script/tile.py | 4 -
src/s_tir/analysis/identify_memcpy.cc | 2 +-
src/s_tir/transform/lower_async_dma.cc | 2 +-
src/target/intrin_rule.cc | 2 +-
src/tirx/ir/buffer.cc | 244 +++++++--------------
src/tirx/ir/buffer_load.cc | 2 +-
src/tirx/ir/stmt.cc | 120 ----------
src/tirx/ir/type.cc | 196 +++++++++++++++++
src/tirx/op/tirx.cc | 1 -
src/tirx/script/builder/frame.cc | 14 --
src/tirx/script/builder/ir.cc | 11 -
src/tirx/script/printer/stmt.cc | 87 +++-----
src/tirx/script/printer/utils.h | 2 +-
src/tirx/transform/lower_intrin.cc | 2 +-
src/tirx/transform/make_packed_api.cc | 2 +-
src/tirx/transform/tvm_ffi_binder.h | 2 +-
src/tirx/transform/vectorize_loop.cc | 2 +-
tests/cpp/sym_simplify_test.cc | 2 +-
tests/cpp/tir_analysis_side_effect.cc | 2 +-
.../operator/tile_primitive/test_dispatcher.py | 16 +-
tests/python/tirx/test_op_namespace_cleanup.py | 1 -
tests/python/tirx/test_parser_printer.py | 54 -----
39 files changed, 506 insertions(+), 785 deletions(-)
diff --git a/docs/tirx/api/tile.rst b/docs/tirx/api/tile.rst
index 485ea34f41..2abf16986b 100644
--- a/docs/tirx/api/tile.rst
+++ b/docs/tirx/api/tile.rst
@@ -35,10 +35,6 @@ the explicit ``Tx.tile`` form consistently. See the
:doc:`programming guide <../tile_primitives>` for the model, primitive catalog,
and dispatch configuration.
-.. automodule:: tvm.tirx.script.tile
- :members: compose_op
- :no-index:
-
Scope namespaces
----------------
diff --git a/docs/tirx/api/tirx.rst b/docs/tirx/api/tirx.rst
index 850ee32324..51c9c1744b 100644
--- a/docs/tirx/api/tirx.rst
+++ b/docs/tirx/api/tirx.rst
@@ -23,6 +23,11 @@ namespace. Layouts, execution scopes, visitors, compilation
helpers, and
tile-dispatch extensions are documented on their focused pages and excluded
here so the same objects are not expanded twice.
+For C++ construction, include ``tvm/tirx/expr.h`` for ``BufferVar``, buffer
+loads, and buffer-region constructors. Include ``tvm/tirx/type.h`` for
+``BufferType``, ``BufferRegionType``, and ``TensorMapType``. Buffer regions
+use the shared ``TensorRegion`` expression from ``tvm/ir/expr.h``.
+
.. automodule:: tvm.tirx
:members:
:imported-members:
diff --git a/docs/tirx/tile_primitives.rst b/docs/tirx/tile_primitives.rst
index b89c1a8380..c3e84e9629 100644
--- a/docs/tirx/tile_primitives.rst
+++ b/docs/tirx/tile_primitives.rst
@@ -93,8 +93,7 @@ introspection.
- ``sum``, ``max``, ``min``
- reduce selected axes, optionally accumulating into the destination
* - Fused and composed
- - ``binary_reduce``, ``unary_reduce``, ``binary_chain``,
``reduce_negate``,
- ``compose_op``
+ - ``binary_reduce``, ``unary_reduce``, ``binary_chain``, ``reduce_negate``
- combine several primitive operations for backends that dispatch them as
one unit
diff --git a/include/tvm/te/operation.h b/include/tvm/te/operation.h
index c59ec67827..33b23f8d83 100644
--- a/include/tvm/te/operation.h
+++ b/include/tvm/te/operation.h
@@ -29,7 +29,7 @@
#include <tvm/ir/prim/expr.h>
#include <tvm/sym/analyzer.h>
#include <tvm/te/tensor.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <string>
diff --git a/include/tvm/tirx/buffer_region.h b/include/tvm/tirx/buffer_region.h
deleted file mode 100644
index edd70443ca..0000000000
--- a/include/tvm/tirx/buffer_region.h
+++ /dev/null
@@ -1,58 +0,0 @@
-/*
- * 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.
- */
-#ifndef TVM_TIRX_BUFFER_REGION_H_
-#define TVM_TIRX_BUFFER_REGION_H_
-
-#include <tvm/ffi/reflection/registry.h>
-#include <tvm/ir/expr.h>
-#include <tvm/tirx/buffer.h>
-
-namespace tvm {
-namespace tirx {
-
-/*! \brief The type of a multi-dimensional buffer region expression. */
-class BufferRegionTypeNode : public TypeNode {
- public:
- static void RegisterReflection() {
- namespace refl = tvm::ffi::reflection;
- refl::ObjectDef<BufferRegionTypeNode>();
- }
-
- TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.BufferRegionType",
BufferRegionTypeNode, TypeNode);
-};
-
-/*! \brief Managed reference to BufferRegionTypeNode. */
-class BufferRegionType : public Type {
- public:
- TVM_DLL BufferRegionType();
-
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(BufferRegionType, Type,
BufferRegionTypeNode);
-};
-
-/*! \brief Construct a region with buffer rank validation and
BufferRegionType. */
-TVM_DLL TensorRegion BufferRegion(BufferVar buffer, ffi::Array<Range> region,
Span span = Span());
-/*! \brief Select the entire buffer. */
-TVM_DLL TensorRegion FullBufferRegion(BufferVar buffer);
-/*! \brief Construct unit or vector-lane ranges from point indices. */
-TVM_DLL TensorRegion BufferRegionFromPoint(BufferVar buffer,
ffi::Array<PrimExpr> indices);
-
-} // namespace tirx
-} // namespace tvm
-
-#endif // TVM_TIRX_BUFFER_REGION_H_
diff --git a/include/tvm/tirx/buffer.h b/include/tvm/tirx/expr.h
similarity index 68%
rename from include/tvm/tirx/buffer.h
rename to include/tvm/tirx/expr.h
index 75a1026760..8d3932c673 100644
--- a/include/tvm/tirx/buffer.h
+++ b/include/tvm/tirx/expr.h
@@ -18,17 +18,17 @@
*/
/*!
- * \file tvm/tirx/buffer.h
- * \brief Symbolic n-dimensional array, to represent a memory buffer.
+ * \file tvm/tirx/expr.h
+ * \brief TIRx buffer expressions and construction helpers.
*/
-#ifndef TVM_TIRX_BUFFER_H_
-#define TVM_TIRX_BUFFER_H_
+#ifndef TVM_TIRX_EXPR_H_
+#define TVM_TIRX_EXPR_H_
#include <tvm/ffi/container/array.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ffi/string.h>
#include <tvm/ir/expr.h>
-#include <tvm/tirx/layout.h>
+#include <tvm/tirx/type.h>
#include <tvm/tirx/var.h>
#include <string>
@@ -36,138 +36,9 @@
namespace tvm {
namespace tirx {
-#ifndef TVM_INDEX_DEFAULT_I64
-#define TVM_INDEX_DEFAULT_I64 1
-#endif
-/*! \brief if TVM_INDEX_DEFAULT_I64 is set, return int64, otherwise return
int32 */
-inline PrimType DefaultIndexPrimType() {
-#if TVM_INDEX_DEFAULT_I64
- static const PrimType default_index_ty = PrimType::Int(64);
-#else
- static const PrimType default_index_ty = PrimType::Int(32);
-#endif
- return default_index_ty;
-}
-
-inline DLDataType DefaultIndexType() {
-#if TVM_INDEX_DEFAULT_I64
- return DLDataType{kDLInt, 64, 1};
-#else
- return DLDataType{kDLInt, 32, 1};
-#endif
-}
-
// forward declare Stmt
class Stmt;
-/*!
- * \brief Structural type of a TIRx buffer variable.
- *
- * A buffer value is an ordinary VarNode whose ExprNode::ty is BufferType.
- * BufferType owns the immutable access contract. The physical pointer is
- * deliberately not stored here; it is obtained with buffer_data(BufferVar)
- * and is bound by the surrounding buffer definition.
- */
-class BufferTypeNode : public TypeNode {
- public:
- /*! \brief dtype in the content of the tensor */
- PrimType dtype = PrimType::Void();
- /*! \brief Storage scope/address space of the buffer. */
- ffi::String storage_scope;
- /*! \brief The type of the buffer prior to flattening
- *
- * This contains the shape as it is accessed by
- * BufferLoad/BufferStore nodes, and used by the low-level code
- * generators.
- */
- ffi::Array<PrimExpr> shape;
- /*!
- * \brief The strides of each dimension
- * This can be an empty array, indicating array is contiguous
- */
- ffi::Array<PrimExpr> strides;
- /*! \brief The offset in terms of number of dtype elements (including lanes)
*/
- PrimExpr elem_offset;
- /*! \brief Alignment requirement of data pointer in bytes. */
- int data_alignment;
- /*!
- * \brief Factor of elem_offset field,
- * elem_offset is guaranteed to be multiple of offset_factor.
- */
- int offset_factor;
- /*! \brief The layout of the buffer */
- ffi::Optional<Layout> layout;
-
- /*! \brief The allocated address of the buffer.
- * The address might be multi-dimensional based on its scope.
- * For example, trn.psum takes 2D address, representing (bank, offset).
- */
- ffi::Array<PrimExpr> allocated_addr;
-
- /*! \brief constructor */
- BufferTypeNode() {}
-
- static void RegisterReflection() {
- namespace refl = tvm::ffi::reflection;
- refl::ObjectDef<BufferTypeNode>()
- .def_ro("dtype", &BufferTypeNode::dtype)
- .def_ro("storage_scope", &BufferTypeNode::storage_scope)
- // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
- .def_ro("shape", &BufferTypeNode::shape,
refl::AttachFieldFlag::SEqHashDefPattern())
- // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
- .def_ro("strides", &BufferTypeNode::strides,
refl::AttachFieldFlag::SEqHashDefPattern())
- // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
- .def_ro("elem_offset", &BufferTypeNode::elem_offset,
- refl::AttachFieldFlag::SEqHashDefPattern())
- .def_ro("data_alignment", &BufferTypeNode::data_alignment)
- .def_ro("offset_factor", &BufferTypeNode::offset_factor)
- .def_ro("layout", &BufferTypeNode::layout)
- .def_ro("allocated_addr", &BufferTypeNode::allocated_addr);
- }
-
- /*! \return preferred index type for this buffer node */
- DLDataType DefaultIndexType() const {
- return shape.size() != 0 ? shape[0].ty()->dtype :
tvm::tirx::DefaultIndexType();
- }
-
- /*! \return primitive element type for compiler-side uses. */
- PrimType ElementType() const { return dtype; }
-
- /*! \return type of the physical pointer projected by buffer_data. */
- PointerType DataPointerType() const { return PointerType(dtype,
storage_scope); }
-
- /*! \brief Determine the offset in the buffer of the given index.
- *
- * Returns the buffer offset, in number of elements of type dtype,
- * without adjusting for number of lanes. (e.g. The number of
- * float16x4 elements in a buffer of type float16x4.)
- *
- * \param index The index to be accessed.
- * \param inner Ignore the elem_offset, return inner offset only
- */
- ffi::Array<PrimExpr> ElemOffset(ffi::Array<PrimExpr> index, bool inner =
false) const;
-
- TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.BufferType", BufferTypeNode,
TypeNode);
-};
-
-/*!
- * \brief Managed reference to BufferTypeNode.
- */
-class BufferType : public Type {
- public:
- TVM_DLL BufferType(ffi::String storage_scope, PrimType dtype,
ffi::Array<PrimExpr> shape,
- ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment,
- int offset_factor, ffi::Optional<Layout> layout =
std::nullopt,
- ffi::Array<PrimExpr> allocated_addr = {}, Span span =
Span());
-
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(BufferType, Type,
BufferTypeNode);
-
- explicit BufferType(ffi::ObjectPtr<BufferTypeNode> n) :
Type(ffi::UnsafeInit{}) {
- TVM_FFI_ICHECK(n != nullptr);
- data_ = std::move(n);
- }
-};
-
/*!
* \brief Checked zero-state view over an ordinary VarNode with BufferType.
*
@@ -370,6 +241,14 @@ TVM_DLL tirx::BufferVar
BufferWithOffsetAlignment(ffi::Array<PrimExpr> shape, Pr
* TensorLoad is required to have a BufferVar source.
*/
TVM_DLL TensorLoad BufferLoad(BufferVar buffer, ffi::Array<PrimExpr> indices,
Span span = Span());
+
+/*! \brief Construct a region with buffer rank validation and
BufferRegionType. */
+TVM_DLL TensorRegion BufferRegion(BufferVar buffer, ffi::Array<Range> region,
Span span = Span());
+/*! \brief Select the entire buffer. */
+TVM_DLL TensorRegion FullBufferRegion(BufferVar buffer);
+/*! \brief Construct unit or vector-lane ranges from point indices. */
+TVM_DLL TensorRegion BufferRegionFromPoint(BufferVar buffer,
ffi::Array<PrimExpr> indices);
+
} // namespace tirx
} // namespace tvm
@@ -412,4 +291,4 @@ struct TypeTraits<tirx::BufferVar> : public
ObjectRefTypeTraitsBase<tirx::Buffer
} // namespace tvm::ffi
-#endif // TVM_TIR_BUFFER_H_
+#endif // TVM_TIRX_EXPR_H_
diff --git a/include/tvm/tirx/function.h b/include/tvm/tirx/function.h
index 67e43a3c02..9010e29510 100644
--- a/include/tvm/tirx/function.h
+++ b/include/tvm/tirx/function.h
@@ -30,7 +30,7 @@
#include <tvm/ir/function.h>
#include <tvm/ir/prim/expr.h>
#include <tvm/runtime/tensor.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/stmt.h>
#include <string>
diff --git a/include/tvm/tirx/script/builder/frame.h
b/include/tvm/tirx/script/builder/frame.h
index 677781956b..0427b9def3 100644
--- a/include/tvm/tirx/script/builder/frame.h
+++ b/include/tvm/tirx/script/builder/frame.h
@@ -639,37 +639,6 @@ class DeclBufferFrame : public TIRFrame {
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(DeclBufferFrame, TIRFrame,
DeclBufferFrameNode);
};
-class ComposeOpFrameNode : public TIRFrameNode {
- public:
- /*! \brief The workspace of the compose op. */
- ffi::Map<ffi::String, tvm::tirx::BufferVar> workspace;
- /*! \brief The config of the compose op. */
- ffi::Map<ffi::String, ffi::Any> config;
- /*! \brief The optional dispatch variant name of the compose op. */
- ffi::Optional<ffi::String> dispatch{std::nullopt};
-
- static void RegisterReflection() {
- namespace refl = tvm::ffi::reflection;
- refl::ObjectDef<ComposeOpFrameNode>()
- .def_ro("workspace", &ComposeOpFrameNode::workspace)
- .def_ro("config", &ComposeOpFrameNode::config)
- .def_ro("dispatch", &ComposeOpFrameNode::dispatch);
- }
- TVM_FFI_DECLARE_OBJECT_INFO_FINAL("script.ir_builder.tirx.ComposeOpFrame",
ComposeOpFrameNode,
- TIRFrameNode);
-
- public:
- void ExitWithScope() final;
-};
-
-class ComposeOpFrame : public TIRFrame {
- public:
- explicit ComposeOpFrame(ffi::ObjectPtr<ComposeOpFrameNode> data) :
TIRFrame(ffi::UnsafeInit{}) {
- TVM_FFI_ICHECK(data != nullptr);
- data_ = std::move(data);
- }
- TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(ComposeOpFrame, TIRFrame,
ComposeOpFrameNode);
-};
class AllocBufferFrameNode : public TIRFrameNode {
public:
/*! \brief The allocated buffer. */
diff --git a/include/tvm/tirx/script/builder/ir.h
b/include/tvm/tirx/script/builder/ir.h
index 48250fbfa4..537c23fb2b 100644
--- a/include/tvm/tirx/script/builder/ir.h
+++ b/include/tvm/tirx/script/builder/ir.h
@@ -475,17 +475,6 @@ LaunchThreadFrame LaunchThread(Var var, PrimExpr extent);
*/
LaunchThreadFrame LaunchThread(ffi::String thread_tag, PrimExpr extent);
-/*!
- * \brief Compose TIRx op.
- * \param workspace The workspace of the compose op.
- * \param config The config of the compose op.
- * \param dispatch The optional dispatch variant name.
- * \return The result ComposeOpFrame.
- */
-ComposeOpFrame ComposeOp(ffi::Map<ffi::String, BufferVar> workspace,
- ffi::Map<ffi::String, ffi::Any> config,
- ffi::Optional<ffi::String> dispatch = std::nullopt);
-
/*!
* \brief Bind a var to thread env.
* \param thread_tag The thread type tag.
diff --git a/include/tvm/tirx/stmt.h b/include/tvm/tirx/stmt.h
index c612f25b0b..212ad12b98 100644
--- a/include/tvm/tirx/stmt.h
+++ b/include/tvm/tirx/stmt.h
@@ -26,9 +26,8 @@
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/expr.h>
-#include <tvm/tirx/buffer.h>
-#include <tvm/tirx/buffer_region.h>
#include <tvm/tirx/exec_scope.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/layout.h>
#include <optional>
diff --git a/include/tvm/tirx/tile_primitive.h
b/include/tvm/tirx/tile_primitive.h
index b0158a25e7..8fbea6becb 100644
--- a/include/tvm/tirx/tile_primitive.h
+++ b/include/tvm/tirx/tile_primitive.h
@@ -337,8 +337,6 @@ TVM_DLL const Op& fma();
TVM_DLL const Op& silu();
-TVM_DLL const Op& compose_op();
-
TVM_DLL const Op& permute_layout();
} // namespace tirx
diff --git a/include/tvm/tirx/type.h b/include/tvm/tirx/type.h
index 20906eeabc..7eb154863e 100644
--- a/include/tvm/tirx/type.h
+++ b/include/tvm/tirx/type.h
@@ -24,10 +24,160 @@
#ifndef TVM_TIRX_TYPE_H_
#define TVM_TIRX_TYPE_H_
+#include <tvm/ir/expr.h>
#include <tvm/ir/type.h>
+#include <tvm/tirx/layout.h>
namespace tvm::tirx {
+#ifndef TVM_INDEX_DEFAULT_I64
+#define TVM_INDEX_DEFAULT_I64 1
+#endif
+/*! \brief if TVM_INDEX_DEFAULT_I64 is set, return int64, otherwise return
int32 */
+inline PrimType DefaultIndexPrimType() {
+#if TVM_INDEX_DEFAULT_I64
+ static const PrimType default_index_ty = PrimType::Int(64);
+#else
+ static const PrimType default_index_ty = PrimType::Int(32);
+#endif
+ return default_index_ty;
+}
+
+inline DLDataType DefaultIndexType() {
+#if TVM_INDEX_DEFAULT_I64
+ return DLDataType{kDLInt, 64, 1};
+#else
+ return DLDataType{kDLInt, 32, 1};
+#endif
+}
+
+/*!
+ * \brief Structural type of a TIRx buffer variable.
+ *
+ * A buffer value is an ordinary VarNode whose ExprNode::ty is BufferType.
+ * BufferType owns the immutable access contract. The physical pointer is
+ * deliberately not stored here; it is obtained with buffer_data(BufferVar)
+ * and is bound by the surrounding buffer definition.
+ */
+class BufferTypeNode : public TypeNode {
+ public:
+ /*! \brief dtype in the content of the tensor */
+ PrimType dtype = PrimType::Void();
+ /*! \brief Storage scope/address space of the buffer. */
+ ffi::String storage_scope;
+ /*! \brief The type of the buffer prior to flattening
+ *
+ * This contains the shape as it is accessed by
+ * BufferLoad/BufferStore nodes, and used by the low-level code
+ * generators.
+ */
+ ffi::Array<PrimExpr> shape;
+ /*!
+ * \brief The strides of each dimension
+ * This can be an empty array, indicating array is contiguous
+ */
+ ffi::Array<PrimExpr> strides;
+ /*! \brief The offset in terms of number of dtype elements (including lanes)
*/
+ PrimExpr elem_offset;
+ /*! \brief Alignment requirement of data pointer in bytes. */
+ int data_alignment;
+ /*!
+ * \brief Factor of elem_offset field,
+ * elem_offset is guaranteed to be multiple of offset_factor.
+ */
+ int offset_factor;
+ /*! \brief The layout of the buffer */
+ ffi::Optional<Layout> layout;
+
+ /*! \brief The allocated address of the buffer.
+ * The address might be multi-dimensional based on its scope.
+ * For example, trn.psum takes 2D address, representing (bank, offset).
+ */
+ ffi::Array<PrimExpr> allocated_addr;
+
+ /*! \brief constructor */
+ BufferTypeNode() {}
+
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<BufferTypeNode>()
+ .def_ro("dtype", &BufferTypeNode::dtype)
+ .def_ro("storage_scope", &BufferTypeNode::storage_scope)
+ // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
+ .def_ro("shape", &BufferTypeNode::shape,
refl::AttachFieldFlag::SEqHashDefPattern())
+ // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
+ .def_ro("strides", &BufferTypeNode::strides,
refl::AttachFieldFlag::SEqHashDefPattern())
+ // TODO(tqchen): use SEqHashDefSimple after the next pypi tvm-ffi
release
+ .def_ro("elem_offset", &BufferTypeNode::elem_offset,
+ refl::AttachFieldFlag::SEqHashDefPattern())
+ .def_ro("data_alignment", &BufferTypeNode::data_alignment)
+ .def_ro("offset_factor", &BufferTypeNode::offset_factor)
+ .def_ro("layout", &BufferTypeNode::layout)
+ .def_ro("allocated_addr", &BufferTypeNode::allocated_addr);
+ }
+
+ /*! \return preferred index type for this buffer node */
+ DLDataType DefaultIndexType() const {
+ return shape.size() != 0 ? shape[0].ty()->dtype :
tvm::tirx::DefaultIndexType();
+ }
+
+ /*! \return primitive element type for compiler-side uses. */
+ PrimType ElementType() const { return dtype; }
+
+ /*! \return type of the physical pointer projected by buffer_data. */
+ PointerType DataPointerType() const { return PointerType(dtype,
storage_scope); }
+
+ /*! \brief Determine the offset in the buffer of the given index.
+ *
+ * Returns the buffer offset, in number of elements of type dtype,
+ * without adjusting for number of lanes. (e.g. The number of
+ * float16x4 elements in a buffer of type float16x4.)
+ *
+ * \param index The index to be accessed.
+ * \param inner Ignore the elem_offset, return inner offset only
+ */
+ ffi::Array<PrimExpr> ElemOffset(ffi::Array<PrimExpr> index, bool inner =
false) const;
+
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.BufferType", BufferTypeNode,
TypeNode);
+};
+
+/*!
+ * \brief Managed reference to BufferTypeNode.
+ */
+class BufferType : public Type {
+ public:
+ TVM_DLL BufferType(ffi::String storage_scope, PrimType dtype,
ffi::Array<PrimExpr> shape,
+ ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment,
+ int offset_factor, ffi::Optional<Layout> layout =
std::nullopt,
+ ffi::Array<PrimExpr> allocated_addr = {}, Span span =
Span());
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(BufferType, Type,
BufferTypeNode);
+
+ explicit BufferType(ffi::ObjectPtr<BufferTypeNode> n) :
Type(ffi::UnsafeInit{}) {
+ TVM_FFI_ICHECK(n != nullptr);
+ data_ = std::move(n);
+ }
+};
+
+/*! \brief The type of a multi-dimensional buffer region expression. */
+class BufferRegionTypeNode : public TypeNode {
+ public:
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<BufferRegionTypeNode>();
+ }
+
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tirx.BufferRegionType",
BufferRegionTypeNode, TypeNode);
+};
+
+/*! \brief Managed reference to BufferRegionTypeNode. */
+class BufferRegionType : public Type {
+ public:
+ TVM_DLL BufferRegionType();
+
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(BufferRegionType, Type,
BufferRegionTypeNode);
+};
+
/*!
* \brief The type of tensor map.
* \sa TensorMapType
diff --git a/python/tvm/backend/trn/tile_primitive/compose_op/__init__.py
b/python/tvm/backend/trn/tile_primitive/compose_op/__init__.py
index b1f28eea18..ebca132fa2 100644
--- a/python/tvm/backend/trn/tile_primitive/compose_op/__init__.py
+++ b/python/tvm/backend/trn/tile_primitive/compose_op/__init__.py
@@ -17,6 +17,5 @@
from .binary_chain import *
from .binary_reduce import *
-from .compose_op import *
from .reduce_negate import *
from .unary_reduce import *
diff --git a/python/tvm/backend/trn/tile_primitive/compose_op/compose_op.py
b/python/tvm/backend/trn/tile_primitive/compose_op/compose_op.py
deleted file mode 100644
index 5fb5a9a201..0000000000
--- a/python/tvm/backend/trn/tile_primitive/compose_op/compose_op.py
+++ /dev/null
@@ -1,47 +0,0 @@
-# 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.
-
-"""Implementation of ComposeOp dispatch."""
-
-from tvm.tirx import PrimFunc, TilePrimitiveCall
-from tvm.tirx.operator.tile_primitive import DispatchContext, predicate,
register_dispatch
-
-
-def compose_op_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc |
None:
- """Generate a TRN schedule for compose operations."""
- raise NotImplementedError(
- "Generic compose_op must be lowered to specific compose ops before
operator-level passes"
- )
-
-
-@register_dispatch(
- "compose_op",
- "trn",
- variant="default",
- priority=10,
- when=[
- predicate(
- "exec_scope",
- lambda op, sctx: (
- sctx.scope_kind == "thread",
- f"unsupported exec_scope {sctx.scope_kind}",
- ),
- )
- ],
-)
-def compose_op_trn_dispatch(op: TilePrimitiveCall, sctx: DispatchContext) ->
PrimFunc:
- return compose_op_trn(op, sctx)
diff --git a/python/tvm/tirx/operator/tile_primitive/ops.py
b/python/tvm/tirx/operator/tile_primitive/ops.py
index 87ff93fef3..9b10f03374 100644
--- a/python/tvm/tirx/operator/tile_primitive/ops.py
+++ b/python/tvm/tirx/operator/tile_primitive/ops.py
@@ -519,31 +519,6 @@ class ReduceNegate(ReduceOp):
reduce_op = ArgProperty(4)
-class ComposeOp(TilePrimitiveCall):
- """Generic operator for composition of multiple operations.
-
- Must be lowered to specific compose operations before operator-level
passes.
- """
-
- # TODO: add a pass to lower generic compose_op to specific compose ops
-
- op = get_tirx_op("compose_op")
-
- @property
- def srcs(self) -> list[Expr]:
- """Get the source expressions (inputs) of the operator."""
- raise NotImplementedError(
- "Generic compose_op must be lowered to specific compose ops before
operator-level passes" # noqa: E501
- )
-
- @property
- def dsts(self) -> list[Expr]:
- """Get the destination expressions (outputs) of the operator."""
- raise NotImplementedError(
- "Generic compose_op must be lowered to specific compose ops before
operator-level passes" # noqa: E501
- )
-
-
class PermuteLayout(TilePrimitiveCall):
"""Move data so the buffer's bytes are arranged under a different layout.
diff --git a/python/tvm/tirx/script/builder/frame.py
b/python/tvm/tirx/script/builder/frame.py
index d36fd5364b..ae3bedcfe6 100644
--- a/python/tvm/tirx/script/builder/frame.py
+++ b/python/tvm/tirx/script/builder/frame.py
@@ -95,10 +95,6 @@ class LaunchThreadFrame(TIRFrame):
return self.iter_var.var
-@_register_object("script.ir_builder.tirx.ComposeOpFrame")
-class ComposeOpFrame(TIRFrame): ...
-
-
@_register_object("script.ir_builder.tirx.AllocBufferFrame")
class AllocBufferFrame(TIRFrame):
def __enter__(self) -> Buffer:
diff --git a/python/tvm/tirx/script/builder/tirx.py
b/python/tvm/tirx/script/builder/tirx.py
index ec2c5eb3e9..17970f6eff 100644
--- a/python/tvm/tirx/script/builder/tirx.py
+++ b/python/tvm/tirx/script/builder/tirx.py
@@ -27,7 +27,7 @@ from tvm.tirx.exec_scope import _SCOPE_KIND_TO_NAME, ExecScope
from tvm.tirx.expr import FloatImm, IntImm
from tvm.tirx.lang.alloc_pool import SMEMPool, TMEMPool
-from . import _ffi_api, frame
+from . import _ffi_api
from .ir import decl_buffer, meta_class
@@ -1315,27 +1315,6 @@ def log2(
)
-def compose_op(
- workspace: dict[str, Buffer] | None = None, dispatch: str | None = None,
**kwargs
-) -> frame.ComposeOpFrame:
- """Compose a TIRx op.
-
- Parameters
- ----------
- workspace : Optional[Dict[str, Buffer]]
- The workspace of the operator
-
- Returns
- -------
- res : frame.ComposeOpFrame
- The result ComposeOpFrame.
- """
- if workspace is None:
- workspace = {}
- config = kwargs or {}
- return _ffi_api.ComposeOp(workspace, config, dispatch) # pylint:
disable=no-member
-
-
@ScopedOp
def binary_reduce(
binary_output: TensorRegion | Buffer,
@@ -1768,7 +1747,6 @@ __all__ = [
"binary_reduce",
"cast",
"cluster",
- "compose_op",
"copy",
"copy_async",
"cta",
diff --git a/python/tvm/tirx/script/tile.py b/python/tvm/tirx/script/tile.py
index ac9a29c478..e891c495ec 100644
--- a/python/tvm/tirx/script/tile.py
+++ b/python/tvm/tirx/script/tile.py
@@ -108,13 +108,9 @@ warpgroup = _builder.ScopeNamespace("warpgroup",
"warpgroup")
warp = _builder.ScopeNamespace("warp", "warp")
thread = _builder.ScopeNamespace("thread", "thread")
-compose_op = _builder.compose_op
-
-
__all__ = [
*_SCOPED_TILE_OP_NAMES,
"cluster",
- "compose_op",
"cta",
"thread",
"warp",
diff --git a/src/s_tir/analysis/identify_memcpy.cc
b/src/s_tir/analysis/identify_memcpy.cc
index 677483c433..76ec928fd5 100644
--- a/src/s_tir/analysis/identify_memcpy.cc
+++ b/src/s_tir/analysis/identify_memcpy.cc
@@ -29,7 +29,7 @@
#include <tvm/sym/int_set.h>
#include <tvm/sym/iter_affine_map.h>
#include <tvm/tirx/analysis.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/stmt.h>
diff --git a/src/s_tir/transform/lower_async_dma.cc
b/src/s_tir/transform/lower_async_dma.cc
index ec1e67285d..f2f1d9fdc5 100644
--- a/src/s_tir/transform/lower_async_dma.cc
+++ b/src/s_tir/transform/lower_async_dma.cc
@@ -30,7 +30,7 @@
#include <tvm/s_tir/transform.h>
#include <tvm/sym/analyzer.h>
#include <tvm/sym/iter_affine_map.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/stmt.h>
#include <optional>
diff --git a/src/target/intrin_rule.cc b/src/target/intrin_rule.cc
index 0d59b1a41f..f3c3cbfffc 100644
--- a/src/target/intrin_rule.cc
+++ b/src/target/intrin_rule.cc
@@ -24,7 +24,7 @@
#include "intrin_rule.h"
#include <tvm/runtime/logging.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/op_attr_types.h>
diff --git a/src/tirx/ir/buffer.cc b/src/tirx/ir/buffer.cc
index f33cefe331..51b8c675ff 100644
--- a/src/tirx/ir/buffer.cc
+++ b/src/tirx/ir/buffer.cc
@@ -20,17 +20,14 @@
/*!
* \file buffer.cc
*/
-#include <tvm/ffi/extra/structural_mutate.h>
-#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/prim/builtin.h>
#include <tvm/ir/prim/expr.h>
-#include <tvm/runtime/device_api.h>
#include <tvm/sym/analyzer.h>
#include <tvm/tirx/analysis.h>
-#include <tvm/tirx/buffer.h>
#include <tvm/tirx/builtin.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/stmt.h>
@@ -47,6 +44,10 @@ using namespace tvm::prim;
namespace {
+using SubscriptSlice = ffi::Array<ffi::Variant<
+ ffi::Tuple<ffi::Optional<PrimExpr>, ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>>,
+ PrimExpr>>;
+
BufferVar RebuildBufferVarFromType(const BufferVar& buffer, BufferType type,
ffi::String name_suffix = "") {
return BufferVar(buffer.name() + name_suffix, std::move(type),
buffer.span());
@@ -111,126 +112,54 @@ ffi::ObjectRef RealizeBufferSubscript(
return BufferRegion(buffer, region, span);
}
-// Structural traversal hooks
-
-TVMFFIAny BufferTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView
value) noexcept {
- // skips: storage_scope, data_alignment, offset_factor
- const BufferTypeNode* self =
- ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value);
- TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->dtype));
- TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->shape));
- // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
- // container; explicit strides retain normal container descent and callback
behavior.
- if (!self->strides.empty()) {
- TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->strides));
- }
-
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->elem_offset));
- TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->layout));
- // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
- // empty container; present addresses retain normal container descent and
callback behavior.
- if (!self->allocated_addr.empty()) {
-
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->allocated_addr));
- }
- return ffi::AnyView(nullptr).CopyToTVMFFIAny();
-}
+ffi::ObjectRef RealizeBufferRegionSubscript(Expr value, SubscriptSlice slice,
Span span) {
+ TensorRegion source = value.as_or_throw<TensorRegion>();
+ TVM_FFI_CHECK_LE(slice.size(), source->region.size(), IndexError)
+ << "Too many indices for a " << source->region.size() << "-dimensional
buffer region";
-TVMFFIAny BufferTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView
value) noexcept {
- // skips: storage_scope, data_alignment, offset_factor
- const BufferTypeNode* self =
- ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value);
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimType>, mapped_dtype,
- mutator->MutateExpected(self->dtype));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_shape,
- mutator->MutateExpected(self->shape));
- // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
- // container; explicit strides retain normal container descent and callback
behavior.
- ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_strides = ffi::Unchanged();
- if (!self->strides.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
descended_strides,
- mutator->MutateExpected(self->strides));
- mapped_strides = std::move(descended_strides);
- }
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_elem_offset,
-
mutator->MutateExpected(self->elem_offset));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Layout>>,
mapped_layout,
- mutator->MutateExpected(self->layout));
- // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
- // empty container; present addresses retain normal container descent and
callback behavior.
- ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_allocated_addr =
ffi::Unchanged();
- if (!self->allocated_addr.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
- descended_allocated_addr,
-
mutator->MutateExpected(self->allocated_addr));
- mapped_allocated_addr = std::move(descended_allocated_addr);
- }
- if (mapped_dtype.UnchangedOrSameAs(self->dtype) &&
mapped_shape.UnchangedOrSameAs(self->shape) &&
- mapped_strides.UnchangedOrSameAs(self->strides) &&
- mapped_elem_offset.UnchangedOrSameAs(self->elem_offset) &&
- mapped_layout.UnchangedOrSameAs(self->layout) &&
- mapped_allocated_addr.UnchangedOrSameAs(self->allocated_addr)) {
- return ffi::Unchanged().CopyToTVMFFIAny();
- }
- ffi::ObjectPtr<BufferTypeNode> copy =
ffi::make_object<BufferTypeNode>(*self);
- copy->dtype =
std::move(mapped_dtype).ValueOrUnchanged(std::move(copy->dtype));
- copy->shape =
std::move(mapped_shape).ValueOrUnchanged(std::move(copy->shape));
- copy->strides =
std::move(mapped_strides).ValueOrUnchanged(std::move(copy->strides));
- copy->elem_offset =
std::move(mapped_elem_offset).ValueOrUnchanged(std::move(copy->elem_offset));
- copy->layout =
std::move(mapped_layout).ValueOrUnchanged(std::move(copy->layout));
- copy->allocated_addr =
-
std::move(mapped_allocated_addr).ValueOrUnchanged(std::move(copy->allocated_addr));
- return
ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy)));
-}
+ bool all_points = slice.size() == source->region.size();
+ for (const auto& item : slice) {
+ if (auto descriptor = item.as<ffi::Tuple<ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>,
+ ffi::Optional<PrimExpr>>>()) {
+ all_points = false;
+ ffi::Optional<PrimExpr> step = descriptor.value().get<2>();
+ TVM_FFI_CHECK(!step.has_value() || is_one(step.value()), ValueError)
+ << "TensorRegion slices with a non-unit step are not supported";
+ }
+ }
-TVMFFIAny BufferTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
- ffi::AnyView value) noexcept {
- // skips: storage_scope, data_alignment, offset_factor
- BufferTypeNode* self = const_cast<BufferTypeNode*>(
- ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimType>, mapped_dtype,
- mutator->MutateExpected(self->dtype,
ffi::InplaceMode::kAllow));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_shape,
- mutator->MutateExpected(self->shape,
ffi::InplaceMode::kAllow));
- // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
- // container; explicit strides retain normal container descent and callback
behavior.
- ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_strides = ffi::Unchanged();
- if (!self->strides.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_strides,
- mutator->MutateExpected(self->strides, ffi::InplaceMode::kAllow));
- mapped_strides = std::move(descended_strides);
- }
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<PrimExpr>, mapped_elem_offset,
- mutator->MutateExpected(self->elem_offset, ffi::InplaceMode::kAllow));
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Optional<Layout>>, mapped_layout,
- mutator->MutateExpected(self->layout, ffi::InplaceMode::kAllow));
- // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
- // empty container; present addresses retain normal container descent and
callback behavior.
- ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_allocated_addr =
ffi::Unchanged();
- if (!self->allocated_addr.empty()) {
- TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
- ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_allocated_addr,
- mutator->MutateExpected(self->allocated_addr,
ffi::InplaceMode::kAllow));
- mapped_allocated_addr = std::move(descended_allocated_addr);
- }
- if (mapped_dtype.UnchangedOrSameAs(self->dtype) &&
mapped_shape.UnchangedOrSameAs(self->shape) &&
- mapped_strides.UnchangedOrSameAs(self->strides) &&
- mapped_elem_offset.UnchangedOrSameAs(self->elem_offset) &&
- mapped_layout.UnchangedOrSameAs(self->layout) &&
- mapped_allocated_addr.UnchangedOrSameAs(self->allocated_addr)) {
- return ffi::Unchanged().CopyToTVMFFIAny();
- }
- if (!mapped_dtype.IsUnchanged()) self->dtype =
std::move(mapped_dtype).ValueUnchecked();
- if (!mapped_shape.IsUnchanged()) self->shape =
std::move(mapped_shape).ValueUnchecked();
- if (!mapped_strides.IsUnchanged()) self->strides =
std::move(mapped_strides).ValueUnchecked();
- if (!mapped_elem_offset.IsUnchanged())
- self->elem_offset = std::move(mapped_elem_offset).ValueUnchecked();
- if (!mapped_layout.IsUnchanged()) self->layout =
std::move(mapped_layout).ValueUnchecked();
- if (!mapped_allocated_addr.IsUnchanged()) {
- self->allocated_addr = std::move(mapped_allocated_addr).ValueUnchecked();
- }
- return ffi::Unchanged().CopyToTVMFFIAny();
+ if (all_points) {
+ ffi::Array<PrimExpr> indices;
+ indices.reserve(slice.size());
+ for (size_t i = 0; i < slice.size(); ++i) {
+ indices.push_back(source->region[i]->min +
slice[i].as<PrimExpr>().value());
+ }
+ return BufferLoad(source->source.as_or_throw<BufferVar>(), indices, span);
+ }
+
+ sym::Analyzer analyzer;
+ ffi::Array<Range> region;
+ region.reserve(source->region.size());
+ for (size_t i = 0; i < slice.size(); ++i) {
+ const Range& old_range = source->region[i];
+ if (auto point = slice[i].as<PrimExpr>()) {
+ PrimExpr new_min = old_range->min + point.value();
+ region.push_back(Range::FromMinExtent(new_min,
IntImm(point.value().ty(), 1)));
+ } else {
+ auto descriptor = slice[i]
+ .as<ffi::Tuple<ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>,
+ ffi::Optional<PrimExpr>>>()
+ .value();
+ PrimExpr start =
descriptor.get<0>().value_or(IntImm(old_range->extent.ty(), 0));
+ PrimExpr stop = descriptor.get<1>().value_or(old_range->extent);
+ region.push_back(
+ Range::FromMinExtent(old_range->min + start, analyzer->Simplify(stop
- start)));
+ }
+ }
+ for (size_t i = slice.size(); i < source->region.size(); ++i) {
+ region.push_back(source->region[i]);
+ }
+ return BufferRegion(source->source.as_or_throw<BufferVar>(), region, span);
}
} // namespace
@@ -238,48 +167,45 @@ TVMFFIAny
BufferTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
using IndexMod = prim::FloorModNode;
using IndexDiv = prim::FloorDivNode;
-BufferType::BufferType(ffi::String storage_scope, PrimType dtype,
ffi::Array<PrimExpr> shape,
- ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment,
- int offset_factor, ffi::Optional<Layout> layout,
- ffi::Array<PrimExpr> allocated_addr, Span span)
- : Type(ffi::UnsafeInit{}) {
- auto n = ffi::make_object<BufferTypeNode>();
- n->dtype = std::move(dtype);
- n->storage_scope = storage_scope.empty() ? ffi::String("global") :
std::move(storage_scope);
- n->shape = std::move(shape);
- n->strides = std::move(strides);
- if (!elem_offset.defined()) {
- elem_offset = IntImm(PrimType(n->DefaultIndexType()), 0);
- }
- n->elem_offset = std::move(elem_offset);
- n->data_alignment =
- data_alignment <= 0 ? static_cast<int>(runtime::kAllocAlignment) :
data_alignment;
- n->offset_factor = offset_factor == 0 ? 1 : offset_factor;
- n->layout = std::move(layout);
- n->allocated_addr = std::move(allocated_addr);
- n->span = std::move(span);
- data_ = std::move(n);
+TensorRegion BufferRegion(BufferVar buffer, ffi::Array<Range> region, Span
span) {
+ TVM_FFI_ICHECK_EQ(buffer->shape.size(), region.size())
+ << "Buffer rank and region dimension mismatch";
+ return TensorRegion(std::move(buffer), std::move(region),
BufferRegionType(), std::move(span));
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ refl::GlobalDef().def("tirx.BufferRegion", [](BufferVar buffer,
ffi::Array<Range> region) {
+ return BufferRegion(buffer, region);
+ });
+}
+
+TensorRegion FullBufferRegion(BufferVar buffer) {
+ ffi::Array<Range> region;
+ for (PrimExpr extent : buffer->shape) {
+ region.push_back(Range::FromMinExtent(0, extent));
+ }
+ return BufferRegion(buffer, region);
+}
+
+TensorRegion BufferRegionFromPoint(BufferVar buffer, ffi::Array<PrimExpr>
indices) {
+ ffi::Array<Range> region;
+ for (const PrimExpr& index : indices) {
+ if (const prim::RampNode* ramp_index = index.as<prim::RampNode>()) {
+ region.push_back(
+ Range::FromMinExtent(ramp_index->base, ramp_index->stride *
ramp_index->lanes));
+ } else {
+ region.push_back(Range::FromMinExtent(index, MakeConst(index.ty(), 1)));
+ }
+ }
+ return BufferRegion(buffer, region);
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
- BufferTypeNode::RegisterReflection();
refl::TypeAttrDef<BufferTypeNode>().def("__subscript_expr_realize__",
RealizeBufferSubscript);
- refl::TypeAttrDef<BufferTypeNode>()
- .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BufferTypeVisit))
- .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BufferTypeMutate))
- .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
- reinterpret_cast<void*>(&BufferTypeMaybeInplaceMutate));
-
- refl::GlobalDef().def(
- "tirx.BufferType",
- [](ffi::String storage_scope, PrimType dtype, ffi::Array<PrimExpr> shape,
- ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment, int offset_factor,
- ffi::Optional<Layout> layout, ffi::Array<PrimExpr> allocated_addr,
Span span) {
- return BufferType(std::move(storage_scope), std::move(dtype),
std::move(shape),
- std::move(strides), std::move(elem_offset),
data_alignment, offset_factor,
- std::move(layout), std::move(allocated_addr),
std::move(span));
- });
+ refl::TypeAttrDef<BufferRegionTypeNode>().def("__subscript_expr_realize__",
+ RealizeBufferRegionSubscript);
}
ffi::Array<PrimExpr> SimplifyArray(sym::AnalyzerObj* ana, ffi::Array<PrimExpr>
array) {
diff --git a/src/tirx/ir/buffer_load.cc b/src/tirx/ir/buffer_load.cc
index 0a04e0cded..9c804b3ce9 100644
--- a/src/tirx/ir/buffer_load.cc
+++ b/src/tirx/ir/buffer_load.cc
@@ -22,7 +22,7 @@
* \brief Buffer-load expression definition.
*/
#include <tvm/ffi/reflection/registry.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
namespace tvm {
namespace tirx {
diff --git a/src/tirx/ir/stmt.cc b/src/tirx/ir/stmt.cc
index 9ed4b8a408..7a2a049d0d 100644
--- a/src/tirx/ir/stmt.cc
+++ b/src/tirx/ir/stmt.cc
@@ -26,7 +26,6 @@
#include <tvm/ffi/function.h>
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/op.h>
-#include <tvm/sym/analyzer.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/op_attr_types.h>
#include <tvm/tirx/stmt.h>
@@ -45,10 +44,6 @@ using namespace tvm::prim;
namespace {
-using SubscriptSlice = ffi::Array<ffi::Variant<
- ffi::Tuple<ffi::Optional<PrimExpr>, ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>>,
- PrimExpr>>;
-
/*!
* \brief Whether an integer literal can be represented exactly by `ty`.
* \note Mirrors the range checks performed by the IntImm constructor.
@@ -659,68 +654,6 @@ TVMFFIAny
BufferStoreMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
return ffi::Unchanged().CopyToTVMFFIAny();
}
-ffi::ObjectRef RealizeBufferRegionSubscript(Expr value, SubscriptSlice slice,
Span span) {
- TensorRegion source = value.as_or_throw<TensorRegion>();
- TVM_FFI_CHECK_LE(slice.size(), source->region.size(), IndexError)
- << "Too many indices for a " << source->region.size() << "-dimensional
buffer region";
-
- bool all_points = slice.size() == source->region.size();
- for (const auto& item : slice) {
- if (auto descriptor = item.as<ffi::Tuple<ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>,
- ffi::Optional<PrimExpr>>>()) {
- all_points = false;
- ffi::Optional<PrimExpr> step = descriptor.value().get<2>();
- TVM_FFI_CHECK(!step.has_value() || is_one(step.value()), ValueError)
- << "TensorRegion slices with a non-unit step are not supported";
- }
- }
-
- if (all_points) {
- ffi::Array<PrimExpr> indices;
- indices.reserve(slice.size());
- for (size_t i = 0; i < slice.size(); ++i) {
- indices.push_back(source->region[i]->min +
slice[i].as<PrimExpr>().value());
- }
- return BufferLoad(source->source.as_or_throw<BufferVar>(), indices, span);
- }
-
- sym::Analyzer analyzer;
- ffi::Array<Range> region;
- region.reserve(source->region.size());
- for (size_t i = 0; i < slice.size(); ++i) {
- const Range& old_range = source->region[i];
- if (auto point = slice[i].as<PrimExpr>()) {
- PrimExpr new_min = old_range->min + point.value();
- region.push_back(Range::FromMinExtent(new_min,
IntImm(point.value().ty(), 1)));
- } else {
- auto descriptor = slice[i]
- .as<ffi::Tuple<ffi::Optional<PrimExpr>,
ffi::Optional<PrimExpr>,
- ffi::Optional<PrimExpr>>>()
- .value();
- PrimExpr start =
descriptor.get<0>().value_or(IntImm(old_range->extent.ty(), 0));
- PrimExpr stop = descriptor.get<1>().value_or(old_range->extent);
- region.push_back(
- Range::FromMinExtent(old_range->min + start, analyzer->Simplify(stop
- start)));
- }
- }
- for (size_t i = slice.size(); i < source->region.size(); ++i) {
- region.push_back(source->region[i]);
- }
- return BufferRegion(source->source.as_or_throw<BufferVar>(), region, span);
-}
-
-TVMFFIAny BufferRegionTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView)
noexcept {
- return ffi::AnyView(nullptr).CopyToTVMFFIAny();
-}
-
-TVMFFIAny BufferRegionTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView)
noexcept {
- return ffi::Unchanged().CopyToTVMFFIAny();
-}
-
-TVMFFIAny BufferRegionTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*,
ffi::AnyView) noexcept {
- return ffi::Unchanged().CopyToTVMFFIAny();
-}
-
TVMFFIAny ScopeIdDefStmtVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView
value) noexcept {
const ScopeIdDefStmtNode* self =
ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
ScopeIdDefStmtNode>(value);
@@ -1275,59 +1208,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
Span span) { return BufferStore(buffer, value,
indices, span); });
}
-// TensorRegion
-BufferRegionType::BufferRegionType() : Type(ffi::UnsafeInit{}) {
- static ffi::ObjectPtr<BufferRegionTypeNode> singleton =
ffi::make_object<BufferRegionTypeNode>();
- data_ = singleton;
-}
-
-TVM_FFI_STATIC_INIT_BLOCK() {
- namespace refl = tvm::ffi::reflection;
- BufferRegionTypeNode::RegisterReflection();
- refl::TypeAttrDef<BufferRegionTypeNode>()
- .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BufferRegionTypeVisit))
- .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BufferRegionTypeMutate))
- .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
- reinterpret_cast<void*>(&BufferRegionTypeMaybeInplaceMutate))
- .def("__subscript_expr_realize__", RealizeBufferRegionSubscript);
-
- refl::GlobalDef().def("tirx.BufferRegionType", []() { return
BufferRegionType(); });
-}
-
-TensorRegion BufferRegion(BufferVar buffer, ffi::Array<Range> region, Span
span) {
- TVM_FFI_ICHECK_EQ(buffer->shape.size(), region.size())
- << "Buffer rank and region dimension mismatch";
- return TensorRegion(std::move(buffer), std::move(region),
BufferRegionType(), std::move(span));
-}
-
-TVM_FFI_STATIC_INIT_BLOCK() {
- namespace refl = tvm::ffi::reflection;
- refl::GlobalDef().def("tirx.BufferRegion", [](BufferVar buffer,
ffi::Array<Range> region) {
- return BufferRegion(buffer, region);
- });
-}
-
-TensorRegion FullBufferRegion(BufferVar buffer) {
- ffi::Array<Range> region;
- for (PrimExpr extent : buffer->shape) {
- region.push_back(Range::FromMinExtent(0, extent));
- }
- return BufferRegion(buffer, region);
-}
-
-TensorRegion BufferRegionFromPoint(BufferVar buffer, ffi::Array<PrimExpr>
indices) {
- ffi::Array<Range> region;
- for (const PrimExpr& index : indices) {
- if (const prim::RampNode* ramp_index = index.as<prim::RampNode>()) {
- region.push_back(
- Range::FromMinExtent(ramp_index->base, ramp_index->stride *
ramp_index->lanes));
- } else {
- region.push_back(Range::FromMinExtent(index, MakeConst(index.ty(), 1)));
- }
- }
- return BufferRegion(buffer, region);
-}
-
// ScopeIdDefStmt
ScopeIdDefStmt::ScopeIdDefStmt(ScopeIdDef def, Span span) {
TVM_FFI_ICHECK(def.defined());
diff --git a/src/tirx/ir/type.cc b/src/tirx/ir/type.cc
index 704c60a40f..2ef4ad53ac 100644
--- a/src/tirx/ir/type.cc
+++ b/src/tirx/ir/type.cc
@@ -24,8 +24,11 @@
#include <tvm/ffi/extra/structural_mutate.h>
#include <tvm/ffi/extra/structural_visit.h>
#include <tvm/ffi/reflection/registry.h>
+#include <tvm/runtime/device_api.h>
#include <tvm/tirx/type.h>
+#include <utility>
+
namespace tvm::tirx {
namespace {
@@ -41,8 +44,201 @@ TVMFFIAny
TensorMapTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*, ffi::AnyVi
return ffi::Unchanged().CopyToTVMFFIAny();
}
+TVMFFIAny BufferTypeVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView
value) noexcept {
+ // skips: storage_scope, data_alignment, offset_factor
+ const BufferTypeNode* self =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value);
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->dtype));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->shape));
+ // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
+ // container; explicit strides retain normal container descent and callback
behavior.
+ if (!self->strides.empty()) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->strides));
+ }
+
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->elem_offset));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->layout));
+ // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
+ // empty container; present addresses retain normal container descent and
callback behavior.
+ if (!self->allocated_addr.empty()) {
+
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->allocated_addr));
+ }
+ return ffi::AnyView(nullptr).CopyToTVMFFIAny();
+}
+
+TVMFFIAny BufferTypeMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView
value) noexcept {
+ // skips: storage_scope, data_alignment, offset_factor
+ const BufferTypeNode* self =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value);
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimType>, mapped_dtype,
+ mutator->MutateExpected(self->dtype));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_shape,
+ mutator->MutateExpected(self->shape));
+ // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
+ // container; explicit strides retain normal container descent and callback
behavior.
+ ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_strides = ffi::Unchanged();
+ if (!self->strides.empty()) {
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
descended_strides,
+ mutator->MutateExpected(self->strides));
+ mapped_strides = std::move(descended_strides);
+ }
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>,
mapped_elem_offset,
+
mutator->MutateExpected(self->elem_offset));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Optional<Layout>>,
mapped_layout,
+ mutator->MutateExpected(self->layout));
+ // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
+ // empty container; present addresses retain normal container descent and
callback behavior.
+ ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_allocated_addr =
ffi::Unchanged();
+ if (!self->allocated_addr.empty()) {
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
+ descended_allocated_addr,
+
mutator->MutateExpected(self->allocated_addr));
+ mapped_allocated_addr = std::move(descended_allocated_addr);
+ }
+ if (mapped_dtype.UnchangedOrSameAs(self->dtype) &&
mapped_shape.UnchangedOrSameAs(self->shape) &&
+ mapped_strides.UnchangedOrSameAs(self->strides) &&
+ mapped_elem_offset.UnchangedOrSameAs(self->elem_offset) &&
+ mapped_layout.UnchangedOrSameAs(self->layout) &&
+ mapped_allocated_addr.UnchangedOrSameAs(self->allocated_addr)) {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+ }
+ ffi::ObjectPtr<BufferTypeNode> copy =
ffi::make_object<BufferTypeNode>(*self);
+ copy->dtype =
std::move(mapped_dtype).ValueOrUnchanged(std::move(copy->dtype));
+ copy->shape =
std::move(mapped_shape).ValueOrUnchanged(std::move(copy->shape));
+ copy->strides =
std::move(mapped_strides).ValueOrUnchanged(std::move(copy->strides));
+ copy->elem_offset =
std::move(mapped_elem_offset).ValueOrUnchanged(std::move(copy->elem_offset));
+ copy->layout =
std::move(mapped_layout).ValueOrUnchanged(std::move(copy->layout));
+ copy->allocated_addr =
+
std::move(mapped_allocated_addr).ValueOrUnchanged(std::move(copy->allocated_addr));
+ return
ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy)));
+}
+
+TVMFFIAny BufferTypeMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
+ ffi::AnyView value) noexcept {
+ // skips: storage_scope, data_alignment, offset_factor
+ BufferTypeNode* self = const_cast<BufferTypeNode*>(
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BufferTypeNode>(value));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimType>, mapped_dtype,
+ mutator->MutateExpected(self->dtype,
ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<ffi::Array<PrimExpr>>,
mapped_shape,
+ mutator->MutateExpected(self->shape,
ffi::InplaceMode::kAllow));
+ // Empty strides denote the common compact layout. Broad callbacks do not
see the empty
+ // container; explicit strides retain normal container descent and callback
behavior.
+ ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_strides = ffi::Unchanged();
+ if (!self->strides.empty()) {
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_strides,
+ mutator->MutateExpected(self->strides, ffi::InplaceMode::kAllow));
+ mapped_strides = std::move(descended_strides);
+ }
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<PrimExpr>, mapped_elem_offset,
+ mutator->MutateExpected(self->elem_offset, ffi::InplaceMode::kAllow));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Optional<Layout>>, mapped_layout,
+ mutator->MutateExpected(self->layout, ffi::InplaceMode::kAllow));
+ // allocated_addr is empty outside specialized storage scopes. Broad
callbacks do not see the
+ // empty container; present addresses retain normal container descent and
callback behavior.
+ ffi::UnchangedOr<ffi::Array<PrimExpr>> mapped_allocated_addr =
ffi::Unchanged();
+ if (!self->allocated_addr.empty()) {
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(
+ ffi::UnchangedOr<ffi::Array<PrimExpr>>, descended_allocated_addr,
+ mutator->MutateExpected(self->allocated_addr,
ffi::InplaceMode::kAllow));
+ mapped_allocated_addr = std::move(descended_allocated_addr);
+ }
+ if (mapped_dtype.UnchangedOrSameAs(self->dtype) &&
mapped_shape.UnchangedOrSameAs(self->shape) &&
+ mapped_strides.UnchangedOrSameAs(self->strides) &&
+ mapped_elem_offset.UnchangedOrSameAs(self->elem_offset) &&
+ mapped_layout.UnchangedOrSameAs(self->layout) &&
+ mapped_allocated_addr.UnchangedOrSameAs(self->allocated_addr)) {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+ }
+ if (!mapped_dtype.IsUnchanged()) self->dtype =
std::move(mapped_dtype).ValueUnchecked();
+ if (!mapped_shape.IsUnchanged()) self->shape =
std::move(mapped_shape).ValueUnchecked();
+ if (!mapped_strides.IsUnchanged()) self->strides =
std::move(mapped_strides).ValueUnchecked();
+ if (!mapped_elem_offset.IsUnchanged())
+ self->elem_offset = std::move(mapped_elem_offset).ValueUnchecked();
+ if (!mapped_layout.IsUnchanged()) self->layout =
std::move(mapped_layout).ValueUnchecked();
+ if (!mapped_allocated_addr.IsUnchanged()) {
+ self->allocated_addr = std::move(mapped_allocated_addr).ValueUnchecked();
+ }
+ return ffi::Unchanged().CopyToTVMFFIAny();
+}
+
+TVMFFIAny BufferRegionTypeVisit(ffi::StructuralVisitorObj*, ffi::AnyView)
noexcept {
+ return ffi::AnyView(nullptr).CopyToTVMFFIAny();
+}
+
+TVMFFIAny BufferRegionTypeMutate(ffi::StructuralMutatorObj*, ffi::AnyView)
noexcept {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+}
+
+TVMFFIAny BufferRegionTypeMaybeInplaceMutate(ffi::StructuralMutatorObj*,
ffi::AnyView) noexcept {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+}
+
} // namespace
+BufferType::BufferType(ffi::String storage_scope, PrimType dtype,
ffi::Array<PrimExpr> shape,
+ ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment,
+ int offset_factor, ffi::Optional<Layout> layout,
+ ffi::Array<PrimExpr> allocated_addr, Span span)
+ : Type(ffi::UnsafeInit{}) {
+ auto n = ffi::make_object<BufferTypeNode>();
+ n->dtype = std::move(dtype);
+ n->storage_scope = storage_scope.empty() ? ffi::String("global") :
std::move(storage_scope);
+ n->shape = std::move(shape);
+ n->strides = std::move(strides);
+ if (!elem_offset.defined()) {
+ elem_offset = IntImm(PrimType(n->DefaultIndexType()), 0);
+ }
+ n->elem_offset = std::move(elem_offset);
+ n->data_alignment =
+ data_alignment <= 0 ? static_cast<int>(runtime::kAllocAlignment) :
data_alignment;
+ n->offset_factor = offset_factor == 0 ? 1 : offset_factor;
+ n->layout = std::move(layout);
+ n->allocated_addr = std::move(allocated_addr);
+ n->span = std::move(span);
+ data_ = std::move(n);
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BufferTypeNode::RegisterReflection();
+ refl::TypeAttrDef<BufferTypeNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BufferTypeVisit))
+ .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BufferTypeMutate))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BufferTypeMaybeInplaceMutate));
+
+ refl::GlobalDef().def(
+ "tirx.BufferType",
+ [](ffi::String storage_scope, PrimType dtype, ffi::Array<PrimExpr> shape,
+ ffi::Array<PrimExpr> strides, PrimExpr elem_offset, int
data_alignment, int offset_factor,
+ ffi::Optional<Layout> layout, ffi::Array<PrimExpr> allocated_addr,
Span span) {
+ return BufferType(std::move(storage_scope), std::move(dtype),
std::move(shape),
+ std::move(strides), std::move(elem_offset),
data_alignment, offset_factor,
+ std::move(layout), std::move(allocated_addr),
std::move(span));
+ });
+}
+
+// TensorRegion
+BufferRegionType::BufferRegionType() : Type(ffi::UnsafeInit{}) {
+ static ffi::ObjectPtr<BufferRegionTypeNode> singleton =
ffi::make_object<BufferRegionTypeNode>();
+ data_ = singleton;
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BufferRegionTypeNode::RegisterReflection();
+ refl::TypeAttrDef<BufferRegionTypeNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BufferRegionTypeVisit))
+ .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BufferRegionTypeMutate))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BufferRegionTypeMaybeInplaceMutate));
+
+ refl::GlobalDef().def("tirx.BufferRegionType", []() { return
BufferRegionType(); });
+}
+
TensorMapType::TensorMapType(Span span) : Type(ffi::UnsafeInit{}) {
ffi::ObjectPtr<TensorMapTypeNode> n = ffi::make_object<TensorMapTypeNode>();
n->span = std::move(span);
diff --git a/src/tirx/op/tirx.cc b/src/tirx/op/tirx.cc
index 21249706f6..02b6a470e8 100644
--- a/src/tirx/op/tirx.cc
+++ b/src/tirx/op/tirx.cc
@@ -169,7 +169,6 @@ TIRX_DEFINE_TILE_OP(cast);
TIRX_DEFINE_TILE_OP(fma);
TIRX_DEFINE_TILE_OP(silu);
TIRX_DEFINE_TILE_OP(permute_layout);
-TIRX_DEFINE_TILE_OP(compose_op);
TIRX_DEFINE_TILE_OP(copy_async);
TIRX_DEFINE_TILE_OP(gemm_async);
diff --git a/src/tirx/script/builder/frame.cc b/src/tirx/script/builder/frame.cc
index 9ae7e19154..05d2e4016b 100644
--- a/src/tirx/script/builder/frame.cc
+++ b/src/tirx/script/builder/frame.cc
@@ -84,7 +84,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
IfFrameNode::RegisterReflection();
ThenFrameNode::RegisterReflection();
ElseFrameNode::RegisterReflection();
- ComposeOpFrameNode::RegisterReflection();
DeclBufferFrameNode::RegisterReflection();
AllocBufferFrameNode::RegisterReflection();
HintFrameNode::RegisterReflection();
@@ -332,19 +331,6 @@ void DeclBufferFrameNode::ExitWithScope() {
}
}
-void ComposeOpFrameNode::ExitWithScope() {
- TIRFrameNode::ExitWithScope();
- ffi::Array<ffi::ObjectRef> ops;
- for (const auto& stmt : stmts) {
- auto op_call = stmt.as<tvm::tirx::TilePrimitiveCallNode>();
- TVM_FFI_ICHECK(op_call) << "ValueError: Only TIRx op calls allowed in
ComposeOp. Violated by "
- << stmt;
- ops.push_back(ffi::GetRef<tvm::tirx::TilePrimitiveCall>(op_call));
- }
- static const Op& compose_op_op = Op::Get("tirx.tile.compose_op");
- AddToParent(tvm::tirx::TilePrimitiveCall(compose_op_op, ops, workspace,
config, dispatch));
-}
-
void AllocBufferFrameNode::ExitWithScope() {
TIRFrameNode::ExitWithScope();
AddToParent(tvm::tirx::SeqStmt::Flatten(tvm::tirx::AllocBuffer(buffer),
AsStmt(stmts)));
diff --git a/src/tirx/script/builder/ir.cc b/src/tirx/script/builder/ir.cc
index f87d6d7cf5..6247683581 100644
--- a/src/tirx/script/builder/ir.cc
+++ b/src/tirx/script/builder/ir.cc
@@ -767,16 +767,6 @@ HintFrame Hint(ffi::String message, ffi::Map<ffi::String,
ffi::Any> attrs) {
return HintFrame(n);
}
-ComposeOpFrame ComposeOp(ffi::Map<ffi::String, BufferVar> workspace,
- ffi::Map<ffi::String, ffi::Any> config,
- ffi::Optional<ffi::String> dispatch) {
- ffi::ObjectPtr<ComposeOpFrameNode> n =
ffi::make_object<ComposeOpFrameNode>();
- n->workspace = workspace;
- n->config = config;
- n->dispatch = dispatch;
- return ComposeOpFrame(n);
-}
-
Var EnvThread(ffi::String thread_tag, PrimType dtype) {
IterVar iter_var(Range{nullptr}, tvm::PrimVar("", dtype),
tvm::tirx::IterVarType::kThreadIndex,
thread_tag);
@@ -1047,7 +1037,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
})
.def("script.ir_builder.tirx.EnvThread", EnvThread)
.def("script.ir_builder.tirx.Hint", Hint)
- .def("script.ir_builder.tirx.ComposeOp", ComposeOp)
.def("script.ir_builder.tirx.BufferStore", BufferStore)
.def("script.ir_builder.tirx.Evaluate", Evaluate)
.def("script.ir_builder.tirx.Ptr", Ptr);
diff --git a/src/tirx/script/printer/stmt.cc b/src/tirx/script/printer/stmt.cc
index 9ce1c69b1f..af9ec33da0 100644
--- a/src/tirx/script/printer/stmt.cc
+++ b/src/tirx/script/printer/stmt.cc
@@ -112,67 +112,36 @@ TVM_FFI_STATIC_INIT_BLOCK() {
}
return TIRx(d, "tile")->Attr(op_name);
};
- if (!op.same_as(tirx::compose_op())) {
- // Trim trailing None args (e.g. optional bias=None, scale=None)
- size_t n_args = op_call->args.size();
- while (n_args > 0 &&
- op_call->args[n_args - 1].type_index() ==
ffi::TypeIndex::kTVMFFINone) {
- --n_args;
- }
- // Detect in-place unary ops: after trimming Nones, if exactly 2 args
- // and args[0]/args[1] refer to the same buffer region, collapse to
1 arg
- bool inplace_unary = false;
- if (n_args == 2) {
- auto dst_opt = op_call->args[0].as<tvm::TensorRegion>();
- auto src_opt = op_call->args[1].as<tvm::TensorRegion>();
- if (dst_opt.has_value() && src_opt.has_value() &&
- dst_opt.value()->source.same_as(src_opt.value()->source) &&
- StructuralEqual()(dst_opt.value()->region,
src_opt.value()->region)) {
- inplace_unary = true;
- }
- }
- ffi::Array<Doc> args;
- for (size_t i = 0; i < n_args; ++i) {
- if (inplace_unary && i == 1) continue; // skip duplicate src
- args.push_back(d->AsDoc<Doc>(op_call->args[i],
p->Attr("args")->ArrayItem(i)));
- }
- ffi::Optional<ExprDoc> disp = std::nullopt;
- if (op_call->dispatch.has_value()) {
- disp = LiteralDoc::Str(op_call->dispatch.value(),
p->Attr("dispatch"));
- }
- return OpCallDoc(scoped_callee(name), args,
- d->AsDoc<DictDoc>(op_call->workspace,
p->Attr("workspace")),
- d->AsDoc<DictDoc>(op_call->config,
p->Attr("config")), disp);
- } else {
- With<TIRFrame> f(d, op_call);
- ffi::Array<tirx::Stmt> stmts;
- for (size_t i = 0, n = op_call->args.size(); i < n; ++i) {
- stmts.push_back(op_call->args[i].as_or_throw<tirx::Stmt>());
- }
- tirx::SeqStmt seq_stmt(stmts);
- AsDocBody(seq_stmt, p->Attr("args"), f->get(), d);
- // Build kwargs: workspace, dispatch, then flatten config
- ffi::Array<ffi::String> kw_keys;
- ffi::Array<ExprDoc> kw_values;
- if (!op_call->workspace.empty()) {
- kw_keys.push_back("workspace");
- kw_values.push_back(d->AsDoc<DictDoc>(op_call->workspace,
p->Attr("workspace")));
- }
- if (op_call->dispatch.has_value()) {
- kw_keys.push_back("dispatch");
- kw_values.push_back(LiteralDoc::Str(op_call->dispatch.value(),
p->Attr("dispatch")));
- }
- using POO = std::pair<ffi::String, ffi::Any>;
- std::vector<POO> items{op_call->config.begin(),
op_call->config.end()};
- std::sort(items.begin(), items.end(),
- [](const POO& a, const POO& b) { return a.first < b.first;
});
- for (const auto& kv : items) {
- kw_keys.push_back(kv.first);
- kw_values.push_back(d->AsDoc<ExprDoc>(kv.second,
p->Attr("config")->MapItem(kv.first)));
+ // Trim trailing None args (e.g. optional bias=None, scale=None)
+ size_t n_args = op_call->args.size();
+ while (n_args > 0 &&
+ op_call->args[n_args - 1].type_index() ==
ffi::TypeIndex::kTVMFFINone) {
+ --n_args;
+ }
+ // Detect in-place unary ops: after trimming Nones, if exactly 2 args
+ // and args[0]/args[1] refer to the same buffer region, collapse to 1
arg
+ bool inplace_unary = false;
+ if (n_args == 2) {
+ auto dst_opt = op_call->args[0].as<tvm::TensorRegion>();
+ auto src_opt = op_call->args[1].as<tvm::TensorRegion>();
+ if (dst_opt.has_value() && src_opt.has_value() &&
+ dst_opt.value()->source.same_as(src_opt.value()->source) &&
+ StructuralEqual()(dst_opt.value()->region,
src_opt.value()->region)) {
+ inplace_unary = true;
}
- return ScopeDoc(std::nullopt, scoped_callee("compose_op")->Call({},
kw_keys, kw_values),
- (*f)->stmts);
}
+ ffi::Array<Doc> args;
+ for (size_t i = 0; i < n_args; ++i) {
+ if (inplace_unary && i == 1) continue; // skip duplicate src
+ args.push_back(d->AsDoc<Doc>(op_call->args[i],
p->Attr("args")->ArrayItem(i)));
+ }
+ ffi::Optional<ExprDoc> disp = std::nullopt;
+ if (op_call->dispatch.has_value()) {
+ disp = LiteralDoc::Str(op_call->dispatch.value(),
p->Attr("dispatch"));
+ }
+ return OpCallDoc(scoped_callee(name), args,
+ d->AsDoc<DictDoc>(op_call->workspace,
p->Attr("workspace")),
+ d->AsDoc<DictDoc>(op_call->config,
p->Attr("config")), disp);
});
}
TVM_SCRIPT_REPR(tirx::TilePrimitiveCallNode, ReprPrintTIR);
diff --git a/src/tirx/script/printer/utils.h b/src/tirx/script/printer/utils.h
index f5b4b4cdd8..172b4e88f9 100644
--- a/src/tirx/script/printer/utils.h
+++ b/src/tirx/script/printer/utils.h
@@ -26,8 +26,8 @@
#include <tvm/s_tir/stmt_functor.h>
#include <tvm/script/printer/ir_docsifier.h>
#include <tvm/tirx/analysis.h>
-#include <tvm/tirx/buffer.h>
#include <tvm/tirx/exec_scope.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/function.h>
#include <tvm/tirx/index_map.h>
#include <tvm/tirx/op.h>
diff --git a/src/tirx/transform/lower_intrin.cc
b/src/tirx/transform/lower_intrin.cc
index c42650bb69..184b85b280 100644
--- a/src/tirx/transform/lower_intrin.cc
+++ b/src/tirx/transform/lower_intrin.cc
@@ -29,8 +29,8 @@
#include <tvm/ir/prim/expr.h>
#include <tvm/runtime/logging.h>
#include <tvm/target/target.h>
-#include <tvm/tirx/buffer.h>
#include <tvm/tirx/builtin.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/transform.h>
diff --git a/src/tirx/transform/make_packed_api.cc
b/src/tirx/transform/make_packed_api.cc
index 70cefb5602..dfc54bd220 100644
--- a/src/tirx/transform/make_packed_api.cc
+++ b/src/tirx/transform/make_packed_api.cc
@@ -30,8 +30,8 @@
#include <tvm/runtime/device_api.h>
#include <tvm/target/target.h>
#include <tvm/tirx/analysis.h>
-#include <tvm/tirx/buffer.h>
#include <tvm/tirx/builtin.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
diff --git a/src/tirx/transform/tvm_ffi_binder.h
b/src/tirx/transform/tvm_ffi_binder.h
index ee0c49a57d..4f51955a12 100644
--- a/src/tirx/transform/tvm_ffi_binder.h
+++ b/src/tirx/transform/tvm_ffi_binder.h
@@ -31,7 +31,7 @@
#include <tvm/ffi/reflection/access_path.h>
#include <tvm/ir/prim/expr.h>
#include <tvm/sym/analyzer.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
#include <tvm/tirx/stmt.h>
#include <string>
diff --git a/src/tirx/transform/vectorize_loop.cc
b/src/tirx/transform/vectorize_loop.cc
index bed3ac8f44..0877454572 100644
--- a/src/tirx/transform/vectorize_loop.cc
+++ b/src/tirx/transform/vectorize_loop.cc
@@ -42,7 +42,7 @@
#include "../../tirx/analysis/check_contains.h"
#include "tvm/ffi/dtype.h"
-#include "tvm/tirx/buffer.h"
+#include "tvm/tirx/expr.h"
namespace tvm {
namespace tirx {
diff --git a/tests/cpp/sym_simplify_test.cc b/tests/cpp/sym_simplify_test.cc
index b8d9d2a81b..6f47e1279f 100644
--- a/tests/cpp/sym_simplify_test.cc
+++ b/tests/cpp/sym_simplify_test.cc
@@ -22,7 +22,7 @@
#include <tvm/runtime/logging.h>
#include <tvm/sym/analyzer.h>
#include <tvm/te/operation.h>
-#include <tvm/tirx/buffer.h>
+#include <tvm/tirx/expr.h>
TEST(Simplify, MinMax) {
tvm::sym::Analyzer ana;
diff --git a/tests/cpp/tir_analysis_side_effect.cc
b/tests/cpp/tir_analysis_side_effect.cc
index 66d2302e8a..06667295c6 100644
--- a/tests/cpp/tir_analysis_side_effect.cc
+++ b/tests/cpp/tir_analysis_side_effect.cc
@@ -23,8 +23,8 @@
#include <tvm/ir/prim/builtin.h>
#include <tvm/runtime/logging.h>
#include <tvm/te/operation.h>
-#include <tvm/tirx/buffer.h>
#include <tvm/tirx/builtin.h>
+#include <tvm/tirx/expr.h>
TEST(SimplePasses, SideEffect) {
using namespace tvm::prim;
diff --git a/tests/python/tirx/operator/tile_primitive/test_dispatcher.py
b/tests/python/tirx/operator/tile_primitive/test_dispatcher.py
index d5bc3b6fe7..ac21a2d7f3 100644
--- a/tests/python/tirx/operator/tile_primitive/test_dispatcher.py
+++ b/tests/python/tirx/operator/tile_primitive/test_dispatcher.py
@@ -100,10 +100,11 @@ def
test_dispatch_forced_variant_missing_table_and_message():
assert "no variant named '__nonexistent__' is registered" in msg
-def test_dispatch_raises_with_aggregated_reasons():
+def test_dispatch_raises_with_aggregated_reasons(monkeypatch):
"""Validate STRICT mode raises aggregated error message with reasons."""
_import_and_register()
from tvm.ir import Op
+ from tvm.tirx.operator.tile_primitive import dispatcher
from tvm.tirx.operator.tile_primitive.dispatcher import run_dispatch
class _OpCall:
@@ -111,8 +112,15 @@ def test_dispatch_raises_with_aggregated_reasons():
self.op = op
self.args = []
- # Use TRN compose_op; variant implementation raises NotImplementedError
- op_call = _OpCall(Op.get("tirx.tile.compose_op"))
+ def failing_impl(op, sctx):
+ raise NotImplementedError("unsupported operation")
+
+ op_call = _OpCall(Op.get("tirx.tile.copy"))
+ monkeypatch.setitem(
+ dispatcher._DISPATCH_TABLE,
+ (op_call.op, "trn"),
+ [dispatcher.DispatchCase("default", 0, [], failing_impl)],
+ )
sctx = _DummySctx(target_kind="trn", exec_scope="thread")
with pytest.raises(RuntimeError) as e:
@@ -120,7 +128,7 @@ def test_dispatch_raises_with_aggregated_reasons():
msg = str(e.value)
print(msg)
- assert "TIRx schedule dispatch failed: op=tirx.tile.compose_op target=trn"
in msg
+ assert "TIRx schedule dispatch failed: op=tirx.tile.copy target=trn" in msg
assert "default" in msg
assert "exception — NotImplementedError" in msg
# opcall content and backtrace should be included inside the table
diff --git a/tests/python/tirx/test_op_namespace_cleanup.py
b/tests/python/tirx/test_op_namespace_cleanup.py
index 78f56aa3d9..ba60b7f74c 100644
--- a/tests/python/tirx/test_op_namespace_cleanup.py
+++ b/tests/python/tirx/test_op_namespace_cleanup.py
@@ -351,7 +351,6 @@ def test_registered_tirx_ops_have_exactly_one_category():
"tirx.add",
"tirx.binary_chain",
"tirx.binary_reduce",
- "tirx.compose_op",
"tirx.copy",
"tirx.copy_async",
"tirx.fdiv",
diff --git a/tests/python/tirx/test_parser_printer.py
b/tests/python/tirx/test_parser_printer.py
index 93fa389888..01d5a1b12d 100644
--- a/tests/python/tirx/test_parser_printer.py
+++ b/tests/python/tirx/test_parser_printer.py
@@ -619,23 +619,6 @@ def test_roundtrip_alloc_under_any_scope():
assert_structural_equal(test, from_source(code))
-def test_roundtrip_compose_op():
- # fmt: off
- @T.prim_func
- def test():
- T.device_entry()
- A = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- B = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- C = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- with Tx.compose_op():
- Tx.add(B, A, T.float32(1))
- Tx.add(C, B, T.float32(1))
- # fmt: on
- code = test.script()
- assert from_source(code).script() == code
- assert_structural_equal(test, from_source(code))
-
-
def test_roundtrip_op_call_workspace():
# fmt: off
@T.prim_func
@@ -651,25 +634,6 @@ def test_roundtrip_op_call_workspace():
assert_structural_equal(test, from_source(code))
-def test_roundtrip_compose_op_call_workspace():
- # fmt: off
- @T.prim_func
- def test():
- T.device_entry()
- A = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- B = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- C = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- psum = T.alloc_buffer([10], "float32", scope="trn.psum")
- intermediate = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- with Tx.compose_op(workspace={"intermediate": intermediate}):
- Tx.add(B, A, T.float32(1))
- Tx.add(C, B, T.float32(1), workspace={"psum": psum})
- # fmt: on
- code = test.script()
- assert from_source(code).script() == code
- assert_structural_equal(test, from_source(code))
-
-
def test_roundtrip_op_call_config():
# fmt: off
@T.prim_func
@@ -684,24 +648,6 @@ def test_roundtrip_op_call_config():
assert_structural_equal(test, from_source(code))
-def test_roundtrip_compose_op_call_config():
- # fmt: off
- @T.prim_func
- def test():
- T.device_entry()
- A = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- B = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- C = T.alloc_buffer([10], "float32", scope="trn.sbuf")
- psum = T.alloc_buffer([10], "float32", scope="trn.psum")
- with Tx.compose_op( schedule="A"):
- Tx.add(B, A, T.float32(1))
- Tx.add(C, B, T.float32(1), workspace={"psum": psum})
- # fmt: on
- code = test.script()
- assert from_source(code).script() == code
- assert_structural_equal(test, from_source(code))
-
-
def test_predicate():
# fmt: off
@T.prim_func