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 7dcc6d965a [REFACTOR][IR] Promote primitive bitwise and shift
operations to nodes (#20392)
7dcc6d965a is described below
commit 7dcc6d965aea86f0a0cd1d7c3c3062ee65b7353e
Author: Tianqi Chen <[email protected]>
AuthorDate: Sat Sep 19 21:31:23 2026 -0400
[REFACTOR][IR] Promote primitive bitwise and shift operations to nodes
(#20392)
Primitive bitwise and shift expressions currently use generic intrinsic
calls, requiring separate matching logic in analysis and lowering. This
change introduces LShift, RShift, BitwiseAnd, BitwiseOr, BitwiseXor, and
BitwiseNot nodes and migrates compiler consumers to typed operands.
Public builders and operators retain their validation and
constant-folding behavior.
The migration covers reflection and traversal, Python and TVMScript,
symbolic analysis including Z3, Relax and TIR transforms, and
C/LLVM/SPIR-V/WebGPU code generation. Existing tests receive
representation-only updates.
Validation: 2,046 existing tests passed, 462 skipped, and 8 expected
failures across C++, symbolic analysis, IR/script, transforms, and
available backend suites. LLVM, CUDA, WebGPU codegen and Z3 compiled
successfully; GPU runtime tests skipped because no GPU was exposed.
SPIR-V source was reviewed but not compiled or executed. Applicable
formatting and lint checks passed.
---
include/tvm/ir/expr_functor.h | 36 +++++
include/tvm/ir/prim/builtin.h | 18 ---
include/tvm/ir/prim/expr.h | 126 ++++++++++++++++
include/tvm/relax/expr_functor.h | 24 +++
python/tvm/ir/prim/__init__.py | 6 +
python/tvm/ir/prim/expr.py | 119 +++++++++++++++
python/tvm/relax/expr_functor.py | 30 ++++
python/tvm/tirx/__init__.py | 1 +
python/tvm/tirx/expr.py | 6 +
python/tvm/tirx/script/builder/ir.py | 12 ++
src/backend/vulkan/codegen/codegen_spirv.cc | 72 +++++----
src/backend/vulkan/codegen/codegen_spirv.h | 6 +
src/backend/webgpu/codegen/codegen_webgpu.cc | 32 ++--
src/backend/webgpu/codegen/codegen_webgpu.h | 2 +
src/ir/expr_functor.cc | 66 ++++++++
src/ir/prim/builtin.cc | 24 ---
src/ir/prim/deep_equal.cc | 11 ++
src/ir/prim/expr.cc | 166 +++++++++++++++++++++
src/ir/prim/op.cc | 12 +-
src/relax/backend/vm/codegen_vm_tir.cc | 6 +
src/relax/ir/expr_functor.cc | 21 +++
src/relax/transform/compute_prim_value.cc | 6 +
src/s_tir/analysis/estimate_flops.cc | 26 ++++
.../feature_extractor/per_store_feature.cc | 6 +
src/s_tir/schedule/ir_comparator.cc | 11 ++
src/s_tir/schedule/ir_comparator.h | 6 +
src/s_tir/transform/rewrite_unsafe_select.cc | 6 +
src/sym/const_int_bound.cc | 33 ++--
src/sym/constraint_helpers.h | 16 +-
src/sym/modular_set.cc | 34 ++---
src/sym/pattern_match.h | 90 ++++++-----
src/sym/rewrite_simplify.cc | 34 +++--
src/sym/rewrite_simplify.h | 2 +
src/sym/z3_prover.cc | 71 ++++-----
src/target/llvm/codegen_llvm.cc | 44 ++++--
src/target/llvm/codegen_llvm.h | 6 +
src/target/source/codegen_c.cc | 52 +++----
src/target/source/codegen_c.h | 62 ++++----
src/tirx/analysis/filter_canonical.cc | 22 +--
src/tirx/analysis/filter_canonical.h | 2 +-
src/tirx/ir/data_type_rewriter.cc | 85 ++++++-----
src/tirx/ir/data_type_rewriter.h | 6 +
src/tirx/ir/specialize.cc | 10 +-
src/tirx/ir/tir_visitor_with_path.cc | 10 ++
src/tirx/ir/tir_visitor_with_path.h | 6 +
src/tirx/op/builtin.cc | 6 -
src/tirx/script/printer/expr.cc | 26 ++++
src/tirx/transform/common_subexpr_elim.cc | 10 ++
src/tirx/transform/tile_primitive_dispatch.cc | 55 +++++--
src/tirx/transform/vectorize_loop.cc | 24 ++-
tests/python/sym/test_sym_rewrite_simplify.py | 2 +-
tests/python/tirx-base/test_tir_nodes.py | 22 +--
tests/python/tirx-base/test_tir_op_types.py | 12 +-
.../tirx-transform/test_tir_transform_vectorize.py | 5 +-
.../operator/tile_primitive/cuda/copy/test_reg.py | 22 ++-
tests/python/tirx/test_layout.py | 22 ++-
56 files changed, 1210 insertions(+), 438 deletions(-)
diff --git a/include/tvm/ir/expr_functor.h b/include/tvm/ir/expr_functor.h
index c28ae17f0b..a3b7671643 100644
--- a/include/tvm/ir/expr_functor.h
+++ b/include/tvm/ir/expr_functor.h
@@ -118,6 +118,24 @@ class ExprFunctor<R(const Expr&, Args...)> {
virtual R Dispatch_(const prim::AddNode* node, Args... args) {
return DispatchDefault_(node, std::forward<Args>(args)...);
}
+ virtual R Dispatch_(const prim::LShiftNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
+ virtual R Dispatch_(const prim::RShiftNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
+ virtual R Dispatch_(const prim::BitwiseAndNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
+ virtual R Dispatch_(const prim::BitwiseOrNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
+ virtual R Dispatch_(const prim::BitwiseXorNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
+ virtual R Dispatch_(const prim::BitwiseNotNode* node, Args... args) {
+ return DispatchDefault_(node, std::forward<Args>(args)...);
+ }
virtual R Dispatch_(const prim::SubNode* node, Args... args) {
return DispatchDefault_(node, std::forward<Args>(args)...);
}
@@ -224,6 +242,12 @@ class ExprFunctor<R(const Expr&, Args...)> {
SetDispatch<TSelf, StringImmNode>(vtable);
SetDispatch<TSelf, prim::CastNode>(vtable);
SetDispatch<TSelf, prim::AddNode>(vtable);
+ SetDispatch<TSelf, prim::LShiftNode>(vtable);
+ SetDispatch<TSelf, prim::RShiftNode>(vtable);
+ SetDispatch<TSelf, prim::BitwiseAndNode>(vtable);
+ SetDispatch<TSelf, prim::BitwiseOrNode>(vtable);
+ SetDispatch<TSelf, prim::BitwiseXorNode>(vtable);
+ SetDispatch<TSelf, prim::BitwiseNotNode>(vtable);
SetDispatch<TSelf, prim::SubNode>(vtable);
SetDispatch<TSelf, prim::MulNode>(vtable);
SetDispatch<TSelf, prim::DivNode>(vtable);
@@ -325,6 +349,12 @@ class TVM_DLL ExprVisitor : public ObjectVisitor {
virtual ffi::Optional<VisitInterrupt> Visit_(const StringImmNode* node);
virtual ffi::Optional<VisitInterrupt> Visit_(const prim::CastNode* node);
virtual ffi::Optional<VisitInterrupt> Visit_(const prim::AddNode* node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::LShiftNode* node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::RShiftNode* node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::BitwiseAndNode*
node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::BitwiseOrNode*
node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::BitwiseXorNode*
node);
+ virtual ffi::Optional<VisitInterrupt> Visit_(const prim::BitwiseNotNode*
node);
virtual ffi::Optional<VisitInterrupt> Visit_(const prim::SubNode* node);
virtual ffi::Optional<VisitInterrupt> Visit_(const prim::MulNode* node);
virtual ffi::Optional<VisitInterrupt> Visit_(const prim::DivNode* node);
@@ -444,6 +474,12 @@ class TVM_DLL ExprMutator : public ObjectMutator {
virtual UnchangedOr<Expr> Mutate_(const StringImmNode* node, InplaceMode
inplace_mode);
virtual UnchangedOr<PrimExpr> Mutate_(const prim::CastNode* node,
InplaceMode inplace_mode);
virtual UnchangedOr<PrimExpr> Mutate_(const prim::AddNode* node, InplaceMode
inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::LShiftNode* node,
InplaceMode inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::RShiftNode* node,
InplaceMode inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseAndNode* node,
InplaceMode inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseOrNode* node,
InplaceMode inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseXorNode* node,
InplaceMode inplace_mode);
+ virtual UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseNotNode* node,
InplaceMode inplace_mode);
virtual UnchangedOr<PrimExpr> Mutate_(const prim::SubNode* node, InplaceMode
inplace_mode);
virtual UnchangedOr<PrimExpr> Mutate_(const prim::MulNode* node, InplaceMode
inplace_mode);
virtual UnchangedOr<PrimExpr> Mutate_(const prim::DivNode* node, InplaceMode
inplace_mode);
diff --git a/include/tvm/ir/prim/builtin.h b/include/tvm/ir/prim/builtin.h
index 87dfece351..0e423e2803 100644
--- a/include/tvm/ir/prim/builtin.h
+++ b/include/tvm/ir/prim/builtin.h
@@ -40,24 +40,6 @@ TVM_DLL const Op& log2();
/*! \brief Count leading zero bits. */
TVM_DLL const Op& clz();
-/*! \brief Left shift. */
-TVM_DLL const Op& shift_left();
-
-/*! \brief Right shift. */
-TVM_DLL const Op& shift_right();
-
-/*! \brief Bitwise and operator. */
-TVM_DLL const Op& bitwise_and();
-
-/*! \brief Bitwise or operator. */
-TVM_DLL const Op& bitwise_or();
-
-/*! \brief Bitwise xor operator. */
-TVM_DLL const Op& bitwise_xor();
-
-/*! \brief Bitwise not operator. */
-TVM_DLL const Op& bitwise_not();
-
/*!
* \brief Same as select, used for unsafe memory access.
*
diff --git a/include/tvm/ir/prim/expr.h b/include/tvm/ir/prim/expr.h
index 361c9860f3..81b17c8ab1 100644
--- a/include/tvm/ir/prim/expr.h
+++ b/include/tvm/ir/prim/expr.h
@@ -114,6 +114,96 @@ class Add : public PrimExpr {
TVM_DEFINE_OBJECT_REF_COW_METHOD(AddNode);
};
+/*! \brief a << b */
+class LShiftNode : public BinaryOpNode<LShiftNode> {
+ public:
+ static constexpr const char* _type_key = "prim.LShift";
+};
+
+/*!
+ * \brief Managed reference to LShiftNode
+ * \sa LShiftNode
+ */
+class LShift : public PrimExpr {
+ public:
+ TVM_DLL LShift(PrimExpr a, PrimExpr b, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(LShift, PrimExpr, LShiftNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(LShiftNode);
+};
+
+/*! \brief a >> b */
+class RShiftNode : public BinaryOpNode<RShiftNode> {
+ public:
+ static constexpr const char* _type_key = "prim.RShift";
+};
+
+/*!
+ * \brief Managed reference to RShiftNode
+ * \sa RShiftNode
+ */
+class RShift : public PrimExpr {
+ public:
+ TVM_DLL RShift(PrimExpr a, PrimExpr b, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(RShift, PrimExpr, RShiftNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(RShiftNode);
+};
+
+/*! \brief a & b */
+class BitwiseAndNode : public BinaryOpNode<BitwiseAndNode> {
+ public:
+ static constexpr const char* _type_key = "prim.BitwiseAnd";
+};
+
+/*!
+ * \brief Managed reference to BitwiseAndNode
+ * \sa BitwiseAndNode
+ */
+class BitwiseAnd : public PrimExpr {
+ public:
+ TVM_DLL BitwiseAnd(PrimExpr a, PrimExpr b, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(BitwiseAnd, PrimExpr,
BitwiseAndNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(BitwiseAndNode);
+};
+
+/*! \brief a | b */
+class BitwiseOrNode : public BinaryOpNode<BitwiseOrNode> {
+ public:
+ static constexpr const char* _type_key = "prim.BitwiseOr";
+};
+
+/*!
+ * \brief Managed reference to BitwiseOrNode
+ * \sa BitwiseOrNode
+ */
+class BitwiseOr : public PrimExpr {
+ public:
+ TVM_DLL BitwiseOr(PrimExpr a, PrimExpr b, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(BitwiseOr, PrimExpr,
BitwiseOrNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(BitwiseOrNode);
+};
+
+/*! \brief a ^ b */
+class BitwiseXorNode : public BinaryOpNode<BitwiseXorNode> {
+ public:
+ static constexpr const char* _type_key = "prim.BitwiseXor";
+};
+
+/*!
+ * \brief Managed reference to BitwiseXorNode
+ * \sa BitwiseXorNode
+ */
+class BitwiseXor : public PrimExpr {
+ public:
+ TVM_DLL BitwiseXor(PrimExpr a, PrimExpr b, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(BitwiseXor, PrimExpr,
BitwiseXorNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(BitwiseXorNode);
+};
+
/*! \brief a - b */
class SubNode : public BinaryOpNode<SubNode> {
public:
@@ -469,6 +559,30 @@ class Not : public PrimExpr {
TVM_DEFINE_OBJECT_REF_COW_METHOD(NotNode);
};
+/*! \brief ~a */
+class BitwiseNotNode : public ExprNode {
+ public:
+ /*! \brief The input operand. */
+ PrimExpr a;
+ static void RegisterReflection() {
+ namespace refl = tvm::ffi::reflection;
+ refl::ObjectDef<BitwiseNotNode>().def_ro("a", &BitwiseNotNode::a);
+ }
+ TVM_FFI_DECLARE_OBJECT_INFO_FINAL("prim.BitwiseNot", BitwiseNotNode,
ExprNode);
+};
+
+/*!
+ * \brief Managed reference to BitwiseNotNode
+ * \sa BitwiseNotNode
+ */
+class BitwiseNot : public PrimExpr {
+ public:
+ TVM_DLL BitwiseNot(PrimExpr a, Span span = Span());
+ TVM_FFI_DEFINE_OBJECT_REF_METHODS_NULLABLE(BitwiseNot, PrimExpr,
BitwiseNotNode);
+ static constexpr bool _type_container_is_exact = true;
+ TVM_DEFINE_OBJECT_REF_COW_METHOD(BitwiseNotNode);
+};
+
/*!
* \brief return true_value if condition is true, otherwise return false_value.
* \note Both true_value and false_value could be evaluated
@@ -586,6 +700,18 @@ inline constexpr bool object_ref_contains_v<PrimExpr,
prim::CastNode> = true;
template <>
inline constexpr bool object_ref_contains_v<PrimExpr, prim::AddNode> = true;
template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::BitwiseNotNode> =
true;
+template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::BitwiseXorNode> =
true;
+template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::BitwiseOrNode> =
true;
+template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::BitwiseAndNode> =
true;
+template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::RShiftNode> = true;
+template <>
+inline constexpr bool object_ref_contains_v<PrimExpr, prim::LShiftNode> = true;
+template <>
inline constexpr bool object_ref_contains_v<PrimExpr, prim::SubNode> = true;
template <>
inline constexpr bool object_ref_contains_v<PrimExpr, prim::MulNode> = true;
diff --git a/include/tvm/relax/expr_functor.h b/include/tvm/relax/expr_functor.h
index 925d84aa1f..6ae7d50610 100644
--- a/include/tvm/relax/expr_functor.h
+++ b/include/tvm/relax/expr_functor.h
@@ -165,6 +165,12 @@ class ExprFunctor<R(const Expr& n, Args...)> {
virtual R VisitExpr_(const StringImmNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const prim::LetNode* op, Args...) EXPR_FUNCTOR_DISABLED;
virtual R VisitExpr_(const prim::AddNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::LShiftNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::RShiftNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::BitwiseAndNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::BitwiseOrNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::BitwiseXorNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
+ virtual R VisitExpr_(const prim::BitwiseNotNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const prim::SubNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const prim::MulNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
virtual R VisitExpr_(const prim::DivNode* op, Args... args)
EXPR_FUNCTOR_DEFAULT;
@@ -210,6 +216,12 @@ class ExprFunctor<R(const Expr& n, Args...)> {
RELAX_EXPR_FUNCTOR_DISPATCH(prim::LetNode);
RELAX_EXPR_FUNCTOR_DISPATCH(TensorLoadNode);
RELAX_EXPR_FUNCTOR_DISPATCH(prim::AddNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::LShiftNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::RShiftNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::BitwiseAndNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::BitwiseOrNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::BitwiseXorNode);
+ RELAX_EXPR_FUNCTOR_DISPATCH(prim::BitwiseNotNode);
RELAX_EXPR_FUNCTOR_DISPATCH(prim::SubNode);
RELAX_EXPR_FUNCTOR_DISPATCH(prim::MulNode);
RELAX_EXPR_FUNCTOR_DISPATCH(prim::DivNode);
@@ -274,6 +286,12 @@ class ExprVisitor : public ExprFunctor<void(const Expr&)> {
void VisitExpr_(const TupleGetItemNode* op) override;
void VisitExpr_(const StringImmNode* op) override;
void VisitExpr_(const prim::AddNode* op) override;
+ void VisitExpr_(const prim::LShiftNode* op) override;
+ void VisitExpr_(const prim::RShiftNode* op) override;
+ void VisitExpr_(const prim::BitwiseAndNode* op) override;
+ void VisitExpr_(const prim::BitwiseOrNode* op) override;
+ void VisitExpr_(const prim::BitwiseXorNode* op) override;
+ void VisitExpr_(const prim::BitwiseNotNode* op) override;
void VisitExpr_(const prim::SubNode* op) override;
void VisitExpr_(const prim::MulNode* op) override;
void VisitExpr_(const prim::DivNode* op) override;
@@ -426,6 +444,12 @@ class ExprMutatorBase : public ExprFunctor<Expr(const
Expr&)> {
Expr VisitExpr_(const TupleGetItemNode* op) override;
Expr VisitExpr_(const StringImmNode* op) override;
Expr VisitExpr_(const prim::AddNode* op) override;
+ Expr VisitExpr_(const prim::LShiftNode* op) override;
+ Expr VisitExpr_(const prim::RShiftNode* op) override;
+ Expr VisitExpr_(const prim::BitwiseAndNode* op) override;
+ Expr VisitExpr_(const prim::BitwiseOrNode* op) override;
+ Expr VisitExpr_(const prim::BitwiseXorNode* op) override;
+ Expr VisitExpr_(const prim::BitwiseNotNode* op) override;
Expr VisitExpr_(const prim::SubNode* op) override;
Expr VisitExpr_(const prim::MulNode* op) override;
Expr VisitExpr_(const prim::DivNode* op) override;
diff --git a/python/tvm/ir/prim/__init__.py b/python/tvm/ir/prim/__init__.py
index 69c7c370e6..3194de88c5 100644
--- a/python/tvm/ir/prim/__init__.py
+++ b/python/tvm/ir/prim/__init__.py
@@ -30,6 +30,10 @@ from .expr import (
Add,
And,
BinaryOpExpr,
+ BitwiseAnd,
+ BitwiseNot,
+ BitwiseOr,
+ BitwiseXor,
Broadcast,
Cast,
CmpExpr,
@@ -40,6 +44,7 @@ from .expr import (
IntImm,
Let,
LogicalExpr,
+ LShift,
Max,
Min,
Mod,
@@ -47,6 +52,7 @@ from .expr import (
Not,
Or,
Ramp,
+ RShift,
Select,
Shuffle,
Sub,
diff --git a/python/tvm/ir/prim/expr.py b/python/tvm/ir/prim/expr.py
index 8f318a7baf..5ae53598f3 100644
--- a/python/tvm/ir/prim/expr.py
+++ b/python/tvm/ir/prim/expr.py
@@ -142,6 +142,125 @@ class Cast(ExprWithOp):
self.__init_handle_by_constructor__(_prim_ffi_api.Cast, dtype, value,
span) # type: ignore
+@tvm_ffi.register_object("prim.LShift")
+class LShift(BinaryOpExpr):
+ """Left shift node.
+
+ Parameters
+ ----------
+ a : Expr
+ The left hand operand.
+
+ b : Expr
+ The right hand operand.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ def __init__(self, a: Expr, b: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.LShift, a, b, span)
# type: ignore
+
+
+@tvm_ffi.register_object("prim.RShift")
+class RShift(BinaryOpExpr):
+ """Right shift node.
+
+ Parameters
+ ----------
+ a : Expr
+ The left hand operand.
+
+ b : Expr
+ The right hand operand.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ def __init__(self, a: Expr, b: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.RShift, a, b, span)
# type: ignore
+
+
+@tvm_ffi.register_object("prim.BitwiseAnd")
+class BitwiseAnd(BinaryOpExpr):
+ """Bitwise and node.
+
+ Parameters
+ ----------
+ a : Expr
+ The left hand operand.
+
+ b : Expr
+ The right hand operand.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ def __init__(self, a: Expr, b: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.BitwiseAnd, a, b,
span) # type: ignore
+
+
+@tvm_ffi.register_object("prim.BitwiseOr")
+class BitwiseOr(BinaryOpExpr):
+ """Bitwise or node.
+
+ Parameters
+ ----------
+ a : Expr
+ The left hand operand.
+
+ b : Expr
+ The right hand operand.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ def __init__(self, a: Expr, b: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.BitwiseOr, a, b,
span) # type: ignore
+
+
+@tvm_ffi.register_object("prim.BitwiseXor")
+class BitwiseXor(BinaryOpExpr):
+ """Bitwise xor node.
+
+ Parameters
+ ----------
+ a : Expr
+ The left hand operand.
+
+ b : Expr
+ The right hand operand.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ def __init__(self, a: Expr, b: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.BitwiseXor, a, b,
span) # type: ignore
+
+
+@tvm_ffi.register_object("prim.BitwiseNot")
+class BitwiseNot(ExprWithOp):
+ """Bitwise not node.
+
+ Parameters
+ ----------
+ a : Expr
+ The input value.
+
+ span : Optional[Span]
+ The location of this expression in the source code.
+ """
+
+ a: Expr
+
+ def __init__(self, a: Expr, span: Span | None = None) -> None:
+ self.__init_handle_by_constructor__(_prim_ffi_api.BitwiseNot, a, span)
# type: ignore
+
+
@tvm_ffi.register_object("prim.Add")
class Add(BinaryOpExpr):
"""Add node.
diff --git a/python/tvm/relax/expr_functor.py b/python/tvm/relax/expr_functor.py
index 6aa353f5a6..ec7a773958 100644
--- a/python/tvm/relax/expr_functor.py
+++ b/python/tvm/relax/expr_functor.py
@@ -165,6 +165,18 @@ class ExprFunctor:
ret = self.visit_floor_div_(expr)
elif isinstance(expr, _tirx.FloorMod):
ret = self.visit_floor_mod_(expr)
+ elif isinstance(expr, _tirx.LShift):
+ ret = self.visit_lshift_(expr)
+ elif isinstance(expr, _tirx.RShift):
+ ret = self.visit_rshift_(expr)
+ elif isinstance(expr, _tirx.BitwiseAnd):
+ ret = self.visit_bitwise_and_(expr)
+ elif isinstance(expr, _tirx.BitwiseOr):
+ ret = self.visit_bitwise_or_(expr)
+ elif isinstance(expr, _tirx.BitwiseXor):
+ ret = self.visit_bitwise_xor_(expr)
+ elif isinstance(expr, _tirx.BitwiseNot):
+ ret = self.visit_bitwise_not_(expr)
elif isinstance(expr, _tirx.Min):
ret = self.visit_min_(expr)
elif isinstance(expr, _tirx.Max):
@@ -273,6 +285,24 @@ class ExprFunctor:
def visit_floor_mod_(self, op: _tirx.FloorMod):
return self.visit_expr_fallback_(op)
+ def visit_lshift_(self, op: _tirx.LShift):
+ return self.visit_expr_fallback_(op)
+
+ def visit_rshift_(self, op: _tirx.RShift):
+ return self.visit_expr_fallback_(op)
+
+ def visit_bitwise_and_(self, op: _tirx.BitwiseAnd):
+ return self.visit_expr_fallback_(op)
+
+ def visit_bitwise_or_(self, op: _tirx.BitwiseOr):
+ return self.visit_expr_fallback_(op)
+
+ def visit_bitwise_xor_(self, op: _tirx.BitwiseXor):
+ return self.visit_expr_fallback_(op)
+
+ def visit_bitwise_not_(self, op: _tirx.BitwiseNot):
+ return self.visit_expr_fallback_(op)
+
def visit_min_(self, op: _tirx.Min):
return self.visit_expr_fallback_(op)
diff --git a/python/tvm/tirx/__init__.py b/python/tvm/tirx/__init__.py
index fc7d831ecf..500e938755 100644
--- a/python/tvm/tirx/__init__.py
+++ b/python/tvm/tirx/__init__.py
@@ -39,6 +39,7 @@ from .type import TensorMapType
from .expr import convert
from .expr import Var, Reduce, FloatImm, IntImm, Cast
from .expr import Add, Sub, Mul, Div, Mod, FloorDiv, FloorMod
+from .expr import LShift, RShift, BitwiseAnd, BitwiseOr, BitwiseXor, BitwiseNot
from .expr import Min, Max, EQ, NE, LT, LE, GT, GE, And, Or, Not
from .expr import Select, BufferLoad, Ramp, Broadcast, Shuffle
from .expr import CallEffectKind, Let, IterVar, CommReducer
diff --git a/python/tvm/tirx/expr.py b/python/tvm/tirx/expr.py
index 7a4f39ba43..4f882fdbbf 100644
--- a/python/tvm/tirx/expr.py
+++ b/python/tvm/tirx/expr.py
@@ -54,6 +54,10 @@ from tvm.ir.prim.expr import ( # noqa: F401
Add,
And,
BinaryOpExpr,
+ BitwiseAnd,
+ BitwiseNot,
+ BitwiseOr,
+ BitwiseXor,
Broadcast,
Cast,
CmpExpr,
@@ -64,6 +68,7 @@ from tvm.ir.prim.expr import ( # noqa: F401
IntImm,
Let,
LogicalExpr,
+ LShift,
Max,
Min,
Mod,
@@ -71,6 +76,7 @@ from tvm.ir.prim.expr import ( # noqa: F401
Not,
Or,
Ramp,
+ RShift,
Select,
Shuffle,
Sub,
diff --git a/python/tvm/tirx/script/builder/ir.py
b/python/tvm/tirx/script/builder/ir.py
index 0bc85ffa65..874e981a27 100644
--- a/python/tvm/tirx/script/builder/ir.py
+++ b/python/tvm/tirx/script/builder/ir.py
@@ -63,6 +63,10 @@ from tvm.tirx.expr import (
NE,
Add,
And,
+ BitwiseAnd,
+ BitwiseNot,
+ BitwiseOr,
+ BitwiseXor,
Broadcast,
BufferLoad,
CallEffectKind,
@@ -74,6 +78,7 @@ from tvm.tirx.expr import (
FloorMod,
IntImm,
IterVar,
+ LShift,
Max,
Min,
Mod,
@@ -82,6 +87,7 @@ from tvm.tirx.expr import (
Or,
Ramp,
Reduce,
+ RShift,
Select,
Shuffle,
Sub,
@@ -3582,6 +3588,12 @@ __all__ = [
"Mod",
"FloorDiv",
"FloorMod",
+ "LShift",
+ "RShift",
+ "BitwiseAnd",
+ "BitwiseOr",
+ "BitwiseXor",
+ "BitwiseNot",
"Min",
"Max",
"EQ",
diff --git a/src/backend/vulkan/codegen/codegen_spirv.cc
b/src/backend/vulkan/codegen/codegen_spirv.cc
index a327a7de82..0c648ac929 100644
--- a/src/backend/vulkan/codegen/codegen_spirv.cc
+++ b/src/backend/vulkan/codegen/codegen_spirv.cc
@@ -337,6 +337,45 @@ spirv::Value CodeGenSPIRV::Dispatch_(const prim::NotNode*
op) {
return builder_->MakeValue(spv::OpLogicalNot, a.stype, a);
}
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::LShiftNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ spirv::Value b = MakeValue(op->b);
+ return builder_->MakeValue(spv::OpShiftLeftLogical, a.stype, a, b);
+}
+
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::BitwiseAndNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ spirv::Value b = MakeValue(op->b);
+ return builder_->MakeValue(spv::OpBitwiseAnd, a.stype, a, b);
+}
+
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::BitwiseOrNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ spirv::Value b = MakeValue(op->b);
+ return builder_->MakeValue(spv::OpBitwiseOr, a.stype, a, b);
+}
+
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::BitwiseXorNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ spirv::Value b = MakeValue(op->b);
+ return builder_->MakeValue(spv::OpBitwiseXor, a.stype, a, b);
+}
+
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::RShiftNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ spirv::Value b = MakeValue(op->b);
+ if (op->a.ty().MatchesCode(DLDataTypeCode::kDLInt)) {
+ return builder_->MakeValue(spv::OpShiftRightArithmetic, a.stype, a, b);
+ } else {
+ return builder_->MakeValue(spv::OpShiftRightLogical, a.stype, a, b);
+ }
+}
+
+spirv::Value CodeGenSPIRV::Dispatch_(const prim::BitwiseNotNode* op) {
+ spirv::Value a = MakeValue(op->a);
+ return builder_->MakeValue(spv::OpNot, a.stype, a);
+}
+
spirv::Value CodeGenSPIRV::Dispatch_(const prim::SelectNode* op) {
return builder_->Select(MakeValue(op->condition), MakeValue(op->true_value),
MakeValue(op->false_value));
@@ -372,39 +411,6 @@ spirv::Value CodeGenSPIRV::Dispatch_(const CallNode* op) {
}
return
builder_->CallGLSL450(builder_->GetSType(op->ty.as_or_throw<PrimType>()),
inst_id,
values);
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- spirv::Value a = MakeValue(op->args[0]);
- spirv::Value b = MakeValue(op->args[1]);
- return builder_->MakeValue(spv::OpBitwiseAnd, a.stype, a, b);
- } else if (op->op.same_as(prim::builtin::bitwise_xor())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- spirv::Value a = MakeValue(op->args[0]);
- spirv::Value b = MakeValue(op->args[1]);
- return builder_->MakeValue(spv::OpBitwiseXor, a.stype, a, b);
- } else if (op->op.same_as(prim::builtin::bitwise_or())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- spirv::Value a = MakeValue(op->args[0]);
- spirv::Value b = MakeValue(op->args[1]);
- return builder_->MakeValue(spv::OpBitwiseOr, a.stype, a, b);
- } else if (op->op.same_as(prim::builtin::bitwise_not())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 1U);
- spirv::Value a = MakeValue(op->args[0]);
- return builder_->MakeValue(spv::OpNot, a.stype, a);
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- spirv::Value a = MakeValue(op->args[0]);
- spirv::Value b = MakeValue(op->args[1]);
- return builder_->MakeValue(spv::OpShiftLeftLogical, a.stype, a, b);
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- spirv::Value a = MakeValue(op->args[0]);
- spirv::Value b = MakeValue(op->args[1]);
- if
(op->args[0].as_or_throw<PrimExpr>().ty().MatchesCode(DLDataTypeCode::kDLInt)) {
- return builder_->MakeValue(spv::OpShiftRightArithmetic, a.stype, a, b);
- } else {
- return builder_->MakeValue(spv::OpShiftRightLogical, a.stype, a, b);
- }
} else if (op->op.same_as(tirx::builtin::reinterpret())) {
return builder_->MakeValue(spv::OpBitcast,
builder_->GetSType(op->ty.as_or_throw<PrimType>()),
MakeValue(op->args[0]));
diff --git a/src/backend/vulkan/codegen/codegen_spirv.h
b/src/backend/vulkan/codegen/codegen_spirv.h
index 6d1e891d31..9531411052 100644
--- a/src/backend/vulkan/codegen/codegen_spirv.h
+++ b/src/backend/vulkan/codegen/codegen_spirv.h
@@ -100,6 +100,12 @@ class CodeGenSPIRV : public
tirx::ExprFunctor<spirv::Value(const Expr&)>,
spirv::Value Dispatch_(const prim::AndNode* op) override;
spirv::Value Dispatch_(const prim::OrNode* op) override;
spirv::Value Dispatch_(const prim::NotNode* op) override;
+ spirv::Value Dispatch_(const prim::LShiftNode* op) override;
+ spirv::Value Dispatch_(const prim::RShiftNode* op) override;
+ spirv::Value Dispatch_(const prim::BitwiseAndNode* op) override;
+ spirv::Value Dispatch_(const prim::BitwiseOrNode* op) override;
+ spirv::Value Dispatch_(const prim::BitwiseXorNode* op) override;
+ spirv::Value Dispatch_(const prim::BitwiseNotNode* op) override;
spirv::Value Dispatch_(const prim::SelectNode* op) override;
spirv::Value Dispatch_(const prim::LetNode* op) override;
spirv::Value Dispatch_(const CallNode* op) override;
diff --git a/src/backend/webgpu/codegen/codegen_webgpu.cc
b/src/backend/webgpu/codegen/codegen_webgpu.cc
index ee9de61624..7d144603bd 100644
--- a/src/backend/webgpu/codegen/codegen_webgpu.cc
+++ b/src/backend/webgpu/codegen/codegen_webgpu.cc
@@ -456,6 +456,24 @@ PrimExpr CodeGenWebGPU::EnforceU32(PrimExpr value) {
return cast(PrimType::UInt(32, value.ty().lanes()), value);
}
+void CodeGenWebGPU::Dispatch_(const prim::LShiftNode* op, std::ostream& os) {
// NOLINT(*)
+ os << '(';
+ this->PrintExpr(op->a, os);
+ os << "<<";
+ // WebGPU requires shift bits to be u32.
+ this->PrintExpr(EnforceU32(op->b), os);
+ os << ')';
+}
+
+void CodeGenWebGPU::Dispatch_(const prim::RShiftNode* op, std::ostream& os) {
// NOLINT(*)
+ os << '(';
+ this->PrintExpr(op->a, os);
+ os << ">>";
+ // WebGPU requires shift bits to be u32.
+ this->PrintExpr(EnforceU32(op->b), os);
+ os << ')';
+}
+
void CodeGenWebGPU::Dispatch_(const CallNode* op, std::ostream& os) { //
NOLINT(*)
TVM_FFI_ICHECK(!op->op.same_as(tirx::builtin::masked_load()))
<< "Predicated buffer load is not supported.";
@@ -468,20 +486,6 @@ void CodeGenWebGPU::Dispatch_(const CallNode* op,
std::ostream& os) { // NOLINT
os << ">(";
this->PrintExpr(op->args[0], os);
os << ")";
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- os << '(';
- this->PrintExpr(op->args[0], os);
- os << ">>";
- // WebGPU requires shift bits to be u32.
- this->PrintExpr(EnforceU32(op->args[1].as_or_throw<PrimExpr>()), os);
- os << ')';
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- os << '(';
- this->PrintExpr(op->args[0], os);
- os << "<<";
- // WebGPU requires shift bits to be u32.
- this->PrintExpr(EnforceU32(op->args[1].as_or_throw<PrimExpr>()), os);
- os << ')';
} else if (op->op.same_as(prim::builtin::if_then_else())) {
// conditional that skips eval if cond evals to false
std::string result = name_supply_->FreshName("condval");
diff --git a/src/backend/webgpu/codegen/codegen_webgpu.h
b/src/backend/webgpu/codegen/codegen_webgpu.h
index e7e4bc1d06..489aba3ac8 100644
--- a/src/backend/webgpu/codegen/codegen_webgpu.h
+++ b/src/backend/webgpu/codegen/codegen_webgpu.h
@@ -70,6 +70,8 @@ class CodeGenWebGPU final : public CodeGenC {
void Dispatch_(const CallNode* op, std::ostream& os) final; //
NOLINT(*)
void Dispatch_(const TensorLoadNode* op, std::ostream& os) final; //
NOLINT(*)
void Dispatch_(const prim::CastNode* op, std::ostream& os) final; //
NOLINT(*)
+ void Dispatch_(const prim::LShiftNode* op, std::ostream& os) final; //
NOLINT(*)
+ void Dispatch_(const prim::RShiftNode* op, std::ostream& os) final; //
NOLINT(*)
void Dispatch_(const prim::SelectNode* op, std::ostream& os) final; //
NOLINT(*)
void Dispatch_(const prim::LetNode* op, std::ostream& os) final; //
NOLINT(*)
void Dispatch_(const FloatImmNode* op, std::ostream& os) final; //
NOLINT(*)
diff --git a/src/ir/expr_functor.cc b/src/ir/expr_functor.cc
index 9d1626e5b3..e076b5c566 100644
--- a/src/ir/expr_functor.cc
+++ b/src/ir/expr_functor.cc
@@ -37,6 +37,12 @@ void ExprVisitor::InitVTable(VTable* vtable) {
SetDispatch<ExprVisitor, StringImmNode>(vtable);
SetDispatch<ExprVisitor, prim::CastNode>(vtable);
SetDispatch<ExprVisitor, prim::AddNode>(vtable);
+ SetDispatch<ExprVisitor, prim::LShiftNode>(vtable);
+ SetDispatch<ExprVisitor, prim::RShiftNode>(vtable);
+ SetDispatch<ExprVisitor, prim::BitwiseAndNode>(vtable);
+ SetDispatch<ExprVisitor, prim::BitwiseOrNode>(vtable);
+ SetDispatch<ExprVisitor, prim::BitwiseXorNode>(vtable);
+ SetDispatch<ExprVisitor, prim::BitwiseNotNode>(vtable);
SetDispatch<ExprVisitor, prim::SubNode>(vtable);
SetDispatch<ExprVisitor, prim::MulNode>(vtable);
SetDispatch<ExprVisitor, prim::DivNode>(vtable);
@@ -153,6 +159,36 @@ ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const
prim::AddNode* node) {
return std::nullopt;
}
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::LShiftNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
+ return std::nullopt;
+}
+
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::RShiftNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
+ return std::nullopt;
+}
+
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::BitwiseAndNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
+ return std::nullopt;
+}
+
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::BitwiseOrNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
+ return std::nullopt;
+}
+
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::BitwiseXorNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
+ return std::nullopt;
+}
+
ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::SubNode* node) {
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->b));
@@ -254,6 +290,11 @@ ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const
prim::NotNode* node) {
return std::nullopt;
}
+ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::BitwiseNotNode*
node) {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->a));
+ return std::nullopt;
+}
+
ffi::Optional<VisitInterrupt> ExprVisitor::Visit_(const prim::SelectNode*
node) {
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->condition));
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(this->Visit(node->true_value));
@@ -305,6 +346,12 @@ void ExprMutator::InitVTable(VTable* vtable) {
SetDispatch<ExprMutator, StringImmNode>(vtable);
SetDispatch<ExprMutator, prim::CastNode>(vtable);
SetDispatch<ExprMutator, prim::AddNode>(vtable);
+ SetDispatch<ExprMutator, prim::LShiftNode>(vtable);
+ SetDispatch<ExprMutator, prim::RShiftNode>(vtable);
+ SetDispatch<ExprMutator, prim::BitwiseAndNode>(vtable);
+ SetDispatch<ExprMutator, prim::BitwiseOrNode>(vtable);
+ SetDispatch<ExprMutator, prim::BitwiseXorNode>(vtable);
+ SetDispatch<ExprMutator, prim::BitwiseNotNode>(vtable);
SetDispatch<ExprMutator, prim::SubNode>(vtable);
SetDispatch<ExprMutator, prim::MulNode>(vtable);
SetDispatch<ExprMutator, prim::DivNode>(vtable);
@@ -520,6 +567,11 @@ UnchangedOr<PrimExpr> ExprMutator::Mutate_(const
prim::CastNode* node, InplaceMo
return PrimExpr(std::move(copy));
\
}
TVM_IR_BINARY_MUTATE_IMPL(Add)
+TVM_IR_BINARY_MUTATE_IMPL(LShift)
+TVM_IR_BINARY_MUTATE_IMPL(RShift)
+TVM_IR_BINARY_MUTATE_IMPL(BitwiseAnd)
+TVM_IR_BINARY_MUTATE_IMPL(BitwiseOr)
+TVM_IR_BINARY_MUTATE_IMPL(BitwiseXor)
TVM_IR_BINARY_MUTATE_IMPL(Sub)
TVM_IR_BINARY_MUTATE_IMPL(Mul)
TVM_IR_BINARY_MUTATE_IMPL(Div)
@@ -551,6 +603,20 @@ UnchangedOr<PrimExpr> ExprMutator::Mutate_(const
prim::NotNode* node, InplaceMod
return PrimExpr(std::move(copy));
}
+UnchangedOr<PrimExpr> ExprMutator::Mutate_(const prim::BitwiseNotNode* node,
+ InplaceMode inplace_mode) {
+ auto a_u = Mutate(node->a, inplace_mode);
+ if (a_u.UnchangedOrSameAs(node->a)) return ffi::Unchanged();
+ if (inplace_mode == InplaceMode::kAllow) {
+ auto* writable = const_cast<prim::BitwiseNotNode*>(node);
+ if (!a_u.IsUnchanged()) writable->a = std::move(a_u).ValueUnchecked();
+ return ffi::Unchanged();
+ }
+ auto copy = ffi::make_object<prim::BitwiseNotNode>(*node);
+ if (!a_u.IsUnchanged()) copy->a = std::move(a_u).ValueUnchecked();
+ return PrimExpr(std::move(copy));
+}
+
UnchangedOr<PrimExpr> ExprMutator::Mutate_(const prim::SelectNode* node,
InplaceMode inplace_mode) {
auto condition_u = Mutate(node->condition, inplace_mode);
auto true_value_u = Mutate(node->true_value, inplace_mode);
diff --git a/src/ir/prim/builtin.cc b/src/ir/prim/builtin.cc
index cc1cc6b6af..5d1def4789 100644
--- a/src/ir/prim/builtin.cc
+++ b/src/ir/prim/builtin.cc
@@ -34,30 +34,6 @@ PRIM_DEFINE_BUILTIN_FUNC(likely)
.set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kExprAnnotation))
.set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(bitwise_and)
- .set_num_inputs(2)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(bitwise_or)
- .set_num_inputs(2)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(bitwise_xor)
- .set_num_inputs(2)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(bitwise_not)
- .set_num_inputs(1)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(shift_left)
- .set_num_inputs(2)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
-PRIM_DEFINE_BUILTIN_FUNC(shift_right)
- .set_num_inputs(2)
- .set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure))
- .set_attr<bool>("TVectorizable", true);
PRIM_DEFINE_BUILTIN_FUNC(if_then_else)
.set_num_inputs(3)
.set_attr<TCallEffectKind>("TCallEffectKind",
static_cast<int64_t>(CallEffectKind::kPure));
diff --git a/src/ir/prim/deep_equal.cc b/src/ir/prim/deep_equal.cc
index ab9302b87f..20cca2a0ff 100644
--- a/src/ir/prim/deep_equal.cc
+++ b/src/ir/prim/deep_equal.cc
@@ -158,6 +158,12 @@ class ExprDeepEqualChecker : private
tvm::ExprFunctor<bool(const Expr&, const Pr
Dispatch(plhs->a, prhs->a);
}
+ bool Dispatch_(const prim::BitwiseNotNode* plhs, const PrimExpr& rhs) final {
+ const auto* prhs = rhs.as<prim::BitwiseNotNode>();
+ return plhs->ty.as_or_throw<PrimType>() ==
prhs->ty.as_or_throw<PrimType>() &&
+ Dispatch(plhs->a, prhs->a);
+ }
+
bool Dispatch_(const prim::SelectNode* plhs, const PrimExpr& rhs) final {
const auto* prhs = rhs.as<prim::SelectNode>();
return plhs->ty.as_or_throw<PrimType>() ==
prhs->ty.as_or_throw<PrimType>() &&
@@ -187,6 +193,11 @@ class ExprDeepEqualChecker : private
tvm::ExprFunctor<bool(const Expr&, const Pr
}
DEFINE_DEEP_EQUAL_BIN_EXPR(prim::AddNode)
+ DEFINE_DEEP_EQUAL_BIN_EXPR(prim::LShiftNode)
+ DEFINE_DEEP_EQUAL_BIN_EXPR(prim::RShiftNode)
+ DEFINE_DEEP_EQUAL_BIN_EXPR(prim::BitwiseAndNode)
+ DEFINE_DEEP_EQUAL_BIN_EXPR(prim::BitwiseOrNode)
+ DEFINE_DEEP_EQUAL_BIN_EXPR(prim::BitwiseXorNode)
DEFINE_DEEP_EQUAL_BIN_EXPR(prim::SubNode)
DEFINE_DEEP_EQUAL_BIN_EXPR(prim::MulNode)
DEFINE_DEEP_EQUAL_BIN_EXPR(prim::DivNode)
diff --git a/src/ir/prim/expr.cc b/src/ir/prim/expr.cc
index 4f3818fcde..93aa5bdf75 100644
--- a/src/ir/prim/expr.cc
+++ b/src/ir/prim/expr.cc
@@ -167,6 +167,42 @@ TVMFFIAny NotMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyView
return ffi::Unchanged().CopyToTVMFFIAny();
}
+TVMFFIAny BitwiseNotVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView
value) noexcept {
+ // skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
+ const BitwiseNotNode* self =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BitwiseNotNode>(value);
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(visitor->VisitExpected(self->a));
+ return ffi::AnyView(nullptr).CopyToTVMFFIAny();
+}
+
+TVMFFIAny BitwiseNotMutate(ffi::StructuralMutatorObj* mutator, ffi::AnyView
value) noexcept {
+ // skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
+ const BitwiseNotNode* self =
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BitwiseNotNode>(value);
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_a_u,
+ mutator->MutateExpected(self->a));
+ if (mapped_a_u.UnchangedOrSameAs(self->a)) {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+ }
+ ffi::ObjectPtr<BitwiseNotNode> copy =
ffi::make_object<BitwiseNotNode>(*self);
+ if (!mapped_a_u.IsUnchanged()) copy->a =
std::move(mapped_a_u).ValueUnchecked();
+ return
ffi::details::AnyUnsafe::MoveAnyToTVMFFIAny(ffi::Any(std::move(copy)));
+}
+
+TVMFFIAny BitwiseNotMaybeInplaceMutate(ffi::StructuralMutatorObj* mutator,
+ ffi::AnyView value) noexcept {
+ // skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
+ BitwiseNotNode* self = const_cast<BitwiseNotNode*>(
+ ffi::details::AnyUnsafe::RawObjectPtrFromAnyViewAfterCheck<const
BitwiseNotNode>(value));
+ TVM_FFI_S_MUTATE_ASSIGN_OR_RETURN(ffi::UnchangedOr<PrimExpr>, mapped_a_u,
+ mutator->MutateExpected(self->a,
ffi::InplaceMode::kAllow));
+ if (mapped_a_u.UnchangedOrSameAs(self->a)) {
+ return ffi::Unchanged().CopyToTVMFFIAny();
+ }
+ if (!mapped_a_u.IsUnchanged()) self->a =
std::move(mapped_a_u).ValueUnchecked();
+ return ffi::Unchanged().CopyToTVMFFIAny();
+}
+
TVMFFIAny SelectVisit(ffi::StructuralVisitorObj* visitor, ffi::AnyView value)
noexcept {
// skips: PrimExpr types are always PrimType and remain unchanged in normal
mutation.
const SelectNode* self =
@@ -314,6 +350,27 @@ TVMFFIAny LetMaybeInplaceMutate(ffi::StructuralMutatorObj*
mutator, ffi::AnyView
data_ = std::move(node);
\
}
+#define TVM_DEFINE_BITWISE_CONSTRUCTOR(Name, AllowBool)
\
+ Name::Name(PrimExpr a, PrimExpr b, Span span) {
\
+ using T = Name::ContainerType;
\
+ TVM_FFI_CHECK(a.defined(), ValueError) << "a is undefined\n";
\
+ TVM_FFI_CHECK(b.defined(), ValueError) << "b is undefined\n";
\
+ const PrimTypeNode* a_ty = GetPrimTypeNode(a);
\
+ const PrimTypeNode* b_ty = GetPrimTypeNode(b);
\
+ TVM_FFI_CHECK(a_ty->dtype == b_ty->dtype, TypeError)
\
+ << "mismatched types. " << a_ty->dtype << " vs. " << b_ty->dtype <<
"\n"; \
+ TVM_FFI_CHECK(a.ty().MatchesCode(DLDataTypeCode::kDLInt,
DLDataTypeCode::kDLUInt) || \
+ (AllowBool &&
a.ty().MatchesCode(DLDataTypeCode::kDLBool)), \
+ TypeError)
\
+ << #Name << " requires integer" << (AllowBool ? " or boolean" : "") <<
" operands"; \
+ ffi::ObjectPtr<T> node = ffi::make_object<T>();
\
+ node->ExprNode::ty = a.get()->ExprNode::ty;
\
+ node->a = std::move(a);
\
+ node->b = std::move(b);
\
+ node->span = std::move(span);
\
+ data_ = std::move(node);
\
+ }
+
#define TVM_DEFINE_CMPOP_CONSTRUCTOR(Name)
\
Name::Name(PrimExpr a, PrimExpr b, Span span) {
\
using T = Name::ContainerType;
\
@@ -380,6 +437,86 @@ TVM_FFI_STATIC_INIT_BLOCK() {
reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<AddNode>));
}
+// LShift
+TVM_DEFINE_BITWISE_CONSTRUCTOR(LShift, false);
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ LShiftNode::RegisterReflection();
+ refl::GlobalDef().def("prim.LShift",
+ [](PrimExpr a, PrimExpr b, Span span) { return
LShift(a, b, span); });
+ refl::TypeAttrDef<LShiftNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BinaryVisit<LShiftNode>))
+ .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BinaryMutate<LShiftNode>))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<LShiftNode>));
+}
+
+// RShift
+TVM_DEFINE_BITWISE_CONSTRUCTOR(RShift, false);
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ RShiftNode::RegisterReflection();
+ refl::GlobalDef().def("prim.RShift",
+ [](PrimExpr a, PrimExpr b, Span span) { return
RShift(a, b, span); });
+ refl::TypeAttrDef<RShiftNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BinaryVisit<RShiftNode>))
+ .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BinaryMutate<RShiftNode>))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<RShiftNode>));
+}
+
+// BitwiseAnd
+TVM_DEFINE_BITWISE_CONSTRUCTOR(BitwiseAnd, true);
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BitwiseAndNode::RegisterReflection();
+ refl::GlobalDef().def("prim.BitwiseAnd",
+ [](PrimExpr a, PrimExpr b, Span span) { return
BitwiseAnd(a, b, span); });
+ refl::TypeAttrDef<BitwiseAndNode>()
+ .attr(refl::type_attr::kStructuralVisit,
+ reinterpret_cast<void*>(&BinaryVisit<BitwiseAndNode>))
+ .attr(refl::type_attr::kStructuralMutate,
+ reinterpret_cast<void*>(&BinaryMutate<BitwiseAndNode>))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+
reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<BitwiseAndNode>));
+}
+
+// BitwiseOr
+TVM_DEFINE_BITWISE_CONSTRUCTOR(BitwiseOr, true);
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BitwiseOrNode::RegisterReflection();
+ refl::GlobalDef().def("prim.BitwiseOr",
+ [](PrimExpr a, PrimExpr b, Span span) { return
BitwiseOr(a, b, span); });
+ refl::TypeAttrDef<BitwiseOrNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BinaryVisit<BitwiseOrNode>))
+ .attr(refl::type_attr::kStructuralMutate,
+ reinterpret_cast<void*>(&BinaryMutate<BitwiseOrNode>))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<BitwiseOrNode>));
+}
+
+// BitwiseXor
+TVM_DEFINE_BITWISE_CONSTRUCTOR(BitwiseXor, true);
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BitwiseXorNode::RegisterReflection();
+ refl::GlobalDef().def("prim.BitwiseXor",
+ [](PrimExpr a, PrimExpr b, Span span) { return
BitwiseXor(a, b, span); });
+ refl::TypeAttrDef<BitwiseXorNode>()
+ .attr(refl::type_attr::kStructuralVisit,
+ reinterpret_cast<void*>(&BinaryVisit<BitwiseXorNode>))
+ .attr(refl::type_attr::kStructuralMutate,
+ reinterpret_cast<void*>(&BinaryMutate<BitwiseXorNode>))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+
reinterpret_cast<void*>(&BinaryMaybeInplaceMutate<BitwiseXorNode>));
+}
+
// Sub
TVM_DEFINE_BINOP_CONSTRUCTOR(Sub);
@@ -677,6 +814,35 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("prim.Not", [](PrimExpr a, Span span) { return Not(a,
span); });
}
+// BitwiseNot
+BitwiseNot::BitwiseNot(PrimExpr a, Span span) {
+ TVM_FFI_CHECK(a.defined(), ValueError) << "a is undefined";
+ PrimType a_ty = a.ty();
+ TVM_FFI_CHECK(a_ty.MatchesCode(DLDataTypeCode::kDLInt,
DLDataTypeCode::kDLUInt) ||
+ a_ty.MatchesCode(DLDataTypeCode::kDLBool),
+ TypeError)
+ << "BitwiseNot requires an integer or boolean operand";
+
+ ffi::ObjectPtr<BitwiseNotNode> node = ffi::make_object<BitwiseNotNode>();
+ node->ExprNode::ty = a_ty;
+ node->a = std::move(a);
+ node->span = std::move(span);
+ data_ = std::move(node);
+}
+
+TVM_FFI_STATIC_INIT_BLOCK() {
+ namespace refl = tvm::ffi::reflection;
+ BitwiseNotNode::RegisterReflection();
+ refl::TypeAttrDef<BitwiseNotNode>()
+ .attr(refl::type_attr::kStructuralVisit,
reinterpret_cast<void*>(&BitwiseNotVisit))
+ .attr(refl::type_attr::kStructuralMutate,
reinterpret_cast<void*>(&BitwiseNotMutate))
+ .attr(refl::type_attr::kStructuralMaybeInplaceMutate,
+ reinterpret_cast<void*>(&BitwiseNotMaybeInplaceMutate));
+
+ refl::GlobalDef().def("prim.BitwiseNot",
+ [](PrimExpr a, Span span) { return BitwiseNot(a,
span); });
+}
+
// Select
Select::Select(PrimExpr condition, PrimExpr true_value, PrimExpr false_value,
Span span) {
TVM_FFI_CHECK(condition.defined(), ValueError) << "condition is undefined";
diff --git a/src/ir/prim/op.cc b/src/ir/prim/op.cc
index 3c4c8e1096..ce6eb93f15 100644
--- a/src/ir/prim/op.cc
+++ b/src/ir/prim/op.cc
@@ -583,7 +583,7 @@ PrimExpr right_shift(PrimExpr a, PrimExpr b, Span span) {
}
});
- return Call(a.ty(), prim::builtin::shift_right(), {a, b}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::RShift(a, b, span);
}
// shift left
@@ -606,7 +606,7 @@ PrimExpr left_shift(PrimExpr a, PrimExpr b, Span span) {
if (pb->value == 0) return a;
}
});
- return Call(a.ty(), prim::builtin::shift_left(), {a, b}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::LShift(a, b, span);
}
// bitwise and
@@ -618,7 +618,7 @@ PrimExpr bitwise_and(PrimExpr a, PrimExpr b, Span span) {
PrimType result_ty = a.ty();
if (pa && pb) return IntImm(result_ty, (pa->value & pb->value), span);
});
- return Call(a.ty(), prim::builtin::bitwise_and(), {a, b}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::BitwiseAnd(a, b, span);
}
// bitwise_or
@@ -630,7 +630,7 @@ PrimExpr bitwise_or(PrimExpr a, PrimExpr b, Span span) {
PrimType result_ty = a.ty();
if (pa && pb) return IntImm(result_ty, (pa->value | pb->value), span);
});
- return Call(a.ty(), prim::builtin::bitwise_or(), {a, b}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::BitwiseOr(a, b, span);
}
// bitwise_xor
@@ -642,7 +642,7 @@ PrimExpr bitwise_xor(PrimExpr a, PrimExpr b, Span span) {
PrimType result_ty = a.ty();
if (pa && pb) return IntImm(result_ty, (pa->value ^ pb->value), span);
});
- return Call(a.ty(), prim::builtin::bitwise_xor(), {a, b}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::BitwiseXor(a, b, span);
}
// bitwise_not
@@ -650,7 +650,7 @@ PrimExpr operator~(PrimExpr a) { return bitwise_neg(a); }
PrimExpr bitwise_neg(PrimExpr a, Span span) {
type_check_int_or_bool_args(a, "~ operator (bitwise NOT)");
- return Call(a.ty(), prim::builtin::bitwise_not(), {a}, {}, {},
span).as_or_throw<PrimExpr>();
+ return prim::BitwiseNot(a, span);
}
TVM_FFI_STATIC_INIT_BLOCK() {
diff --git a/src/relax/backend/vm/codegen_vm_tir.cc
b/src/relax/backend/vm/codegen_vm_tir.cc
index 3877a70842..d2571780eb 100644
--- a/src/relax/backend/vm/codegen_vm_tir.cc
+++ b/src/relax/backend/vm/codegen_vm_tir.cc
@@ -305,6 +305,12 @@ class CodeGenVMTIR : public
ExprFunctor<ffi::Optional<Expr>(const Expr&)> {
VM_TIR_PRIM_EXPR(TensorLoadNode);
VM_TIR_PRIM_EXPR(prim::AddNode);
+ VM_TIR_PRIM_EXPR(prim::LShiftNode);
+ VM_TIR_PRIM_EXPR(prim::RShiftNode);
+ VM_TIR_PRIM_EXPR(prim::BitwiseAndNode);
+ VM_TIR_PRIM_EXPR(prim::BitwiseOrNode);
+ VM_TIR_PRIM_EXPR(prim::BitwiseXorNode);
+ VM_TIR_PRIM_EXPR(prim::BitwiseNotNode);
VM_TIR_PRIM_EXPR(prim::SubNode);
VM_TIR_PRIM_EXPR(prim::MulNode);
VM_TIR_PRIM_EXPR(prim::DivNode);
diff --git a/src/relax/ir/expr_functor.cc b/src/relax/ir/expr_functor.cc
index f675edf034..f7ec9d1816 100644
--- a/src/relax/ir/expr_functor.cc
+++ b/src/relax/ir/expr_functor.cc
@@ -192,6 +192,11 @@ void ExprVisitor::VisitExpr_(const TensorLoadNode* op) {
}
RELAX_VISIT_TIRX_BINOP(AddNode);
+RELAX_VISIT_TIRX_BINOP(LShiftNode);
+RELAX_VISIT_TIRX_BINOP(RShiftNode);
+RELAX_VISIT_TIRX_BINOP(BitwiseAndNode);
+RELAX_VISIT_TIRX_BINOP(BitwiseOrNode);
+RELAX_VISIT_TIRX_BINOP(BitwiseXorNode);
RELAX_VISIT_TIRX_BINOP(SubNode);
RELAX_VISIT_TIRX_BINOP(MulNode);
RELAX_VISIT_TIRX_BINOP(DivNode);
@@ -223,6 +228,12 @@ void ExprVisitor::VisitExpr_(const prim::NotNode* op) {
VisitExprDepTypeFieldIfNeeded(this, op->ty);
}
+void ExprVisitor::VisitExpr_(const prim::BitwiseNotNode* op) {
+ this->VisitSpan(op->span);
+ this->VisitExpr(op->a);
+ VisitExprDepTypeFieldIfNeeded(this, op->ty);
+}
+
void ExprVisitor::VisitExpr_(const prim::SelectNode* op) {
this->VisitSpan(op->span);
this->VisitExpr(op->condition);
@@ -545,6 +556,11 @@ Expr ExprMutatorBase::VisitExpr_(const TensorLoadNode* op)
{
}
RELAX_MUTATE_TIRX_BINOP(Add);
+RELAX_MUTATE_TIRX_BINOP(LShift);
+RELAX_MUTATE_TIRX_BINOP(RShift);
+RELAX_MUTATE_TIRX_BINOP(BitwiseAnd);
+RELAX_MUTATE_TIRX_BINOP(BitwiseOr);
+RELAX_MUTATE_TIRX_BINOP(BitwiseXor);
RELAX_MUTATE_TIRX_BINOP(Sub);
RELAX_MUTATE_TIRX_BINOP(Mul);
RELAX_MUTATE_TIRX_BINOP(Div);
@@ -576,6 +592,11 @@ Expr ExprMutatorBase::VisitExpr_(const prim::NotNode* op) {
return a.same_as(op->a) ? ffi::GetRef<Expr>(op) : Expr(prim::Not(a,
op->span));
}
+Expr ExprMutatorBase::VisitExpr_(const prim::BitwiseNotNode* op) {
+ PrimExpr a = this->VisitExpr(op->a).as_or_throw<PrimExpr>();
+ return a.same_as(op->a) ? ffi::GetRef<Expr>(op) : Expr(prim::BitwiseNot(a,
op->span));
+}
+
Expr ExprMutatorBase::VisitExpr_(const prim::SelectNode* op) {
PrimExpr condition = this->VisitExpr(op->condition).as_or_throw<PrimExpr>();
PrimExpr true_value =
this->VisitExpr(op->true_value).as_or_throw<PrimExpr>();
diff --git a/src/relax/transform/compute_prim_value.cc
b/src/relax/transform/compute_prim_value.cc
index 2a6160ba42..b18aa3506e 100644
--- a/src/relax/transform/compute_prim_value.cc
+++ b/src/relax/transform/compute_prim_value.cc
@@ -64,6 +64,12 @@ class PrimExprComputeInjector : public ExprMutator {
RELAX_LIFT_PRIM_EXPR(TensorLoadNode);
RELAX_LIFT_PRIM_EXPR(prim::AddNode);
+ RELAX_LIFT_PRIM_EXPR(prim::LShiftNode);
+ RELAX_LIFT_PRIM_EXPR(prim::RShiftNode);
+ RELAX_LIFT_PRIM_EXPR(prim::BitwiseAndNode);
+ RELAX_LIFT_PRIM_EXPR(prim::BitwiseOrNode);
+ RELAX_LIFT_PRIM_EXPR(prim::BitwiseXorNode);
+ RELAX_LIFT_PRIM_EXPR(prim::BitwiseNotNode);
RELAX_LIFT_PRIM_EXPR(prim::SubNode);
RELAX_LIFT_PRIM_EXPR(prim::MulNode);
RELAX_LIFT_PRIM_EXPR(prim::DivNode);
diff --git a/src/s_tir/analysis/estimate_flops.cc
b/src/s_tir/analysis/estimate_flops.cc
index c8112dc80f..2411dc7e41 100644
--- a/src/s_tir/analysis/estimate_flops.cc
+++ b/src/s_tir/analysis/estimate_flops.cc
@@ -127,6 +127,32 @@ class FlopEstimator : private
tirx::ExprFunctor<TResult(const Expr& n)>,
}
}
+ TResult Dispatch_(const prim::LShiftNode* op) final {
+ TResult result = Dispatch(op->a);
+ result += Dispatch(op->b);
+ return result;
+ }
+ TResult Dispatch_(const prim::RShiftNode* op) final {
+ TResult result = Dispatch(op->a);
+ result += Dispatch(op->b);
+ return result;
+ }
+ TResult Dispatch_(const prim::BitwiseAndNode* op) final {
+ TResult result = Dispatch(op->a);
+ result += Dispatch(op->b);
+ return result;
+ }
+ TResult Dispatch_(const prim::BitwiseOrNode* op) final {
+ TResult result = Dispatch(op->a);
+ result += Dispatch(op->b);
+ return result;
+ }
+ TResult Dispatch_(const prim::BitwiseXorNode* op) final {
+ TResult result = Dispatch(op->a);
+ result += Dispatch(op->b);
+ return result;
+ }
+ TResult Dispatch_(const prim::BitwiseNotNode* op) final { return
Dispatch(op->a); }
TResult Dispatch_(const prim::NotNode* op) override { return
Dispatch(op->a); }
TResult Dispatch_(const prim::AndNode* op) final {
TResult result = Dispatch(op->a);
diff --git a/src/s_tir/meta_schedule/feature_extractor/per_store_feature.cc
b/src/s_tir/meta_schedule/feature_extractor/per_store_feature.cc
index ee79c7ae37..939a136eec 100644
--- a/src/s_tir/meta_schedule/feature_extractor/per_store_feature.cc
+++ b/src/s_tir/meta_schedule/feature_extractor/per_store_feature.cc
@@ -578,6 +578,12 @@ Feature::ArithOps::ArithOps(const BufferStoreNode* store,
int64_t prod_loop_exte
TVM_FEATURE_SIMPLE(OrNode, bool_op);
TVM_FEATURE_SIMPLE(NotNode, bool_op);
TVM_FEATURE_SIMPLE(SelectNode, select_op);
+ TVM_FEATURE_SIMPLE(prim::LShiftNode, int_math_func);
+ TVM_FEATURE_SIMPLE(prim::RShiftNode, int_math_func);
+ TVM_FEATURE_SIMPLE(prim::BitwiseAndNode, int_math_func);
+ TVM_FEATURE_SIMPLE(prim::BitwiseOrNode, int_math_func);
+ TVM_FEATURE_SIMPLE(prim::BitwiseXorNode, int_math_func);
+ TVM_FEATURE_SIMPLE(prim::BitwiseNotNode, int_math_func);
TVM_FEATURE_BINARY(AddNode, float_add_sub, int_add_sub);
TVM_FEATURE_BINARY(SubNode, float_add_sub, int_add_sub);
TVM_FEATURE_BINARY(MulNode, float_mul, int_mul);
diff --git a/src/s_tir/schedule/ir_comparator.cc
b/src/s_tir/schedule/ir_comparator.cc
index e82df9bb13..fa0f4c074e 100644
--- a/src/s_tir/schedule/ir_comparator.cc
+++ b/src/s_tir/schedule/ir_comparator.cc
@@ -339,6 +339,17 @@ TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(MaxNode);
TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(FloorDivNode);
TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(FloorModNode);
+TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(prim::LShiftNode);
+TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(prim::RShiftNode);
+TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(prim::BitwiseAndNode);
+TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(prim::BitwiseOrNode);
+TVM_DECLARE_TENSORIZE_COMPARATOR_BINOP(prim::BitwiseXorNode);
+
+bool TensorizeComparator::Dispatch_(const prim::BitwiseNotNode* op, const
PrimExpr& other) {
+ const auto* rhs = other.as<prim::BitwiseNotNode>();
+ return Dispatch(op->a, rhs->a);
+}
+
bool TensorizeComparator::Dispatch_(const IntImmNode* op, const PrimExpr&
other) {
const auto* rhs = other.as<IntImmNode>();
if (op->value != rhs->value) {
diff --git a/src/s_tir/schedule/ir_comparator.h
b/src/s_tir/schedule/ir_comparator.h
index 2cf0444f2d..4558bcf8e3 100644
--- a/src/s_tir/schedule/ir_comparator.h
+++ b/src/s_tir/schedule/ir_comparator.h
@@ -57,6 +57,12 @@ class TensorizeComparator : public ExprComparator, public
StmtComparator {
bool Dispatch_(const SBlockRealizeNode* op, const Stmt& other) override;
bool Dispatch_(const SBlockNode* op, const Stmt& other) override;
+ bool Dispatch_(const prim::LShiftNode* op, const PrimExpr& other) override;
+ bool Dispatch_(const prim::RShiftNode* op, const PrimExpr& other) override;
+ bool Dispatch_(const prim::BitwiseAndNode* op, const PrimExpr& other)
override;
+ bool Dispatch_(const prim::BitwiseOrNode* op, const PrimExpr& other)
override;
+ bool Dispatch_(const prim::BitwiseXorNode* op, const PrimExpr& other)
override;
+ bool Dispatch_(const prim::BitwiseNotNode* op, const PrimExpr& other)
override;
bool Dispatch_(const AddNode* op, const PrimExpr& other) override;
bool Dispatch_(const SubNode* op, const PrimExpr& other) override;
bool Dispatch_(const MulNode* op, const PrimExpr& other) override;
diff --git a/src/s_tir/transform/rewrite_unsafe_select.cc
b/src/s_tir/transform/rewrite_unsafe_select.cc
index 87f84f0f77..264c021598 100644
--- a/src/s_tir/transform/rewrite_unsafe_select.cc
+++ b/src/s_tir/transform/rewrite_unsafe_select.cc
@@ -89,6 +89,12 @@ class UnsafeExprDetector : public
tirx::ExprFunctor<bool(const Expr& n)> {
bool Dispatch_(const prim::GENode* op) final { return BinaryOp(op); }
bool Dispatch_(const prim::AndNode* op) final { return BinaryOp(op); }
bool Dispatch_(const prim::OrNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::LShiftNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::RShiftNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::BitwiseAndNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::BitwiseOrNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::BitwiseXorNode* op) final { return BinaryOp(op); }
+ bool Dispatch_(const prim::BitwiseNotNode* op) final { return
Dispatch(op->a); }
bool Dispatch_(const prim::NotNode* op) final { return Dispatch(op->a); }
bool Dispatch_(const prim::LetNode* op) final {
return Dispatch(op->body) || Dispatch(op->value);
diff --git a/src/sym/const_int_bound.cc b/src/sym/const_int_bound.cc
index 2eba4ffd35..d4ffd4bb5b 100644
--- a/src/sym/const_int_bound.cc
+++ b/src/sym/const_int_bound.cc
@@ -473,21 +473,6 @@ class ConstIntBoundAnalyzer::Impl
return Union(a, b);
}
- Entry Dispatch_(const CallNode* op) final {
- // only special handle >> and & which can be
- // used for index calculation.
-
- if (op->op.same_as(prim::builtin::shift_right())) {
- return VisitRightShift(op);
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- return VisitLeftShift(op);
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- return VisitBitwiseAnd(op);
- } else {
- return Everything(op->ty.as_or_throw<PrimType>());
- }
- }
-
Entry Dispatch_(const VarNode* op) final {
Var v = ffi::GetRef<Var>(op);
auto it = var_map_.find(v);
@@ -498,9 +483,9 @@ class ConstIntBoundAnalyzer::Impl
}
}
- Entry VisitLeftShift(const CallNode* op) {
- Entry a = Dispatch(op->args[0].as_or_throw<PrimExpr>());
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::LShiftNode* op) final {
+ Entry a = Dispatch(op->a);
+ Entry b = Dispatch(op->b);
if (a.min_value < 0 || b.min_value < 0) {
// If either operand can negative, we may run into undefined
@@ -512,15 +497,15 @@ class ConstIntBoundAnalyzer::Impl
return BinaryOpBoundary(a, b, InfAwareLeftShift);
}
- Entry VisitRightShift(const CallNode* op) {
- Entry a = Dispatch(op->args[0].as_or_throw<PrimExpr>());
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::RShiftNode* op) final {
+ Entry a = Dispatch(op->a);
+ Entry b = Dispatch(op->b);
return BinaryOpBoundary(a, b, InfAwareRightShift);
}
- Entry VisitBitwiseAnd(const CallNode* op) {
- Entry a = Dispatch(op->args[0].as_or_throw<PrimExpr>());
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::BitwiseAndNode* op) final {
+ Entry a = Dispatch(op->a);
+ Entry b = Dispatch(op->b);
// handle positive index case.
if (a.min_value >= 0 && b.min_value >= 0) {
return MakeBound(0, std::min(a.max_value, b.max_value));
diff --git a/src/sym/constraint_helpers.h b/src/sym/constraint_helpers.h
index f405b87e02..fbfcd90dc3 100644
--- a/src/sym/constraint_helpers.h
+++ b/src/sym/constraint_helpers.h
@@ -109,16 +109,12 @@ inline void CollectDerivedConstraintFacts(const PrimExpr&
condition, std::vector
CollectDerivedConstraintFacts(and_node->b, out);
return;
}
- if (const auto* call = condition.as<CallNode>()) {
- if (call->op.same_as(prim::builtin::bitwise_and()) && call->args.size() ==
2) {
- PrimExpr lhs = call->args[0].as_or_throw<PrimExpr>();
- PrimExpr rhs = call->args[1].as_or_throw<PrimExpr>();
- if (lhs.ty().MatchesElementType(DLDataTypeCode::kDLBool, 8) &&
- rhs.ty().MatchesElementType(DLDataTypeCode::kDLBool, 8)) {
- CollectDerivedConstraintFacts(lhs, out);
- CollectDerivedConstraintFacts(rhs, out);
- return;
- }
+ if (const auto* and_node = condition.as<prim::BitwiseAndNode>()) {
+ if (and_node->a.ty().MatchesElementType(DLDataTypeCode::kDLBool, 8) &&
+ and_node->b.ty().MatchesElementType(DLDataTypeCode::kDLBool, 8)) {
+ CollectDerivedConstraintFacts(and_node->a, out);
+ CollectDerivedConstraintFacts(and_node->b, out);
+ return;
}
}
if (const auto* eq = condition.as<prim::EQNode>()) {
diff --git a/src/sym/modular_set.cc b/src/sym/modular_set.cc
index 6cdf859996..16b129ad82 100644
--- a/src/sym/modular_set.cc
+++ b/src/sym/modular_set.cc
@@ -261,20 +261,6 @@ class ModularSetAnalyzer::Impl : public
tvm::ExprFunctor<ModularSetAnalyzer::Ent
return Everything();
}
- Entry Dispatch_(const CallNode* op) final {
- // only special handle >> which can be
- // used for index calculation.
- if (op->op.same_as(prim::builtin::shift_right())) {
- return VisitRightShift(op);
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- return VisitBitwiseAnd(op);
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- return VisitLeftShift(op);
- } else {
- return Everything();
- }
- }
-
Entry Dispatch_(const VarNode* op) final {
Var v = ffi::GetRef<Var>(op);
auto it = var_map_.find(v);
@@ -285,32 +271,30 @@ class ModularSetAnalyzer::Impl : public
tvm::ExprFunctor<ModularSetAnalyzer::Ent
}
}
- Entry VisitLeftShift(const CallNode* op) {
- Entry a = Dispatch(op->args[0].as_or_throw<PrimExpr>());
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::LShiftNode* op) final {
+ Entry a = Dispatch(op->a);
+ Entry b = Dispatch(op->b);
if (b.is_const()) {
return Entry(a.coeff << b.base, a.base << b.base);
}
return Everything();
}
- Entry VisitRightShift(const CallNode* op) {
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::RShiftNode* op) final {
+ Entry b = Dispatch(op->b);
// a c x / c -> a x
if (b.is_const()) {
- return DivByConst(op->args[0].as_or_throw<PrimExpr>(),
static_cast<int64_t>(1) << b.base,
- true);
+ return DivByConst(op->a, static_cast<int64_t>(1) << b.base, true);
}
return Everything();
}
- Entry VisitBitwiseAnd(const CallNode* op) {
- Entry b = Dispatch(op->args[1].as_or_throw<PrimExpr>());
+ Entry Dispatch_(const prim::BitwiseAndNode* op) final {
+ Entry b = Dispatch(op->b);
if (b.is_const()) {
int shift;
if (is_const_power_of_two_integer(IntImm::Int32(b.base + 1), &shift)) {
- return ModByConst(op->args[0].as_or_throw<PrimExpr>(),
static_cast<int64_t>(1) << shift,
- true);
+ return ModByConst(op->a, static_cast<int64_t>(1) << shift, true);
}
}
return Everything();
diff --git a/src/sym/pattern_match.h b/src/sym/pattern_match.h
index 9f6c9913db..420830a448 100644
--- a/src/sym/pattern_match.h
+++ b/src/sym/pattern_match.h
@@ -330,9 +330,10 @@ class PConst : public Pattern<PConst<T>> {
* \tparam OpType The AST noderef type.
* \tparam TA The pattern type of the first operand.
* \tparam TB The pattern type of the second operand.
+ * \tparam kFoldConstants Whether evaluation folds constant operands.
*/
-template <typename OpType, typename TA, typename TB>
-class PBinaryExpr : public Pattern<PBinaryExpr<OpType, TA, TB>> {
+template <typename OpType, typename TA, typename TB, bool kFoldConstants =
true>
+class PBinaryExpr : public Pattern<PBinaryExpr<OpType, TA, TB,
kFoldConstants>> {
public:
PBinaryExpr(const TA& a, const TB& b) : a_(a), b_(b) {}
@@ -355,7 +356,9 @@ class PBinaryExpr : public Pattern<PBinaryExpr<OpType, TA,
TB>> {
PrimExpr Eval() const {
PrimExpr lhs = a_.Eval();
PrimExpr rhs = b_.Eval();
- if (auto ret = TryConstFold<OpType>(lhs, rhs)) return ret.value();
+ if constexpr (kFoldConstants) {
+ if (auto ret = TryConstFold<OpType>(lhs, rhs)) return ret.value();
+ }
return OpType(lhs, rhs);
}
@@ -423,6 +426,22 @@ TVM_PATTERN_BINARY_OP(truncmod, prim::Mod);
TVM_PATTERN_BINARY_OP(floordiv, prim::FloorDiv);
TVM_PATTERN_BINARY_OP(floormod, prim::FloorMod);
+// Bitwise patterns preserve nonfolding construction during evaluation.
+#define TVM_PATTERN_BITWISE_OP(FuncName, NodeName)
\
+ template <typename TA, typename TB>
\
+ inline PBinaryExpr<NodeName, TA, TB, false> FuncName(const Pattern<TA>& a,
\
+ const Pattern<TB>& b) {
\
+ return PBinaryExpr<NodeName, TA, TB, false>(a.derived(), b.derived());
\
+ }
+
+TVM_PATTERN_BITWISE_OP(operator<<, prim::LShift);
+TVM_PATTERN_BITWISE_OP(operator>>, prim::RShift);
+TVM_PATTERN_BITWISE_OP(operator&, prim::BitwiseAnd);
+TVM_PATTERN_BITWISE_OP(operator|, prim::BitwiseOr);
+TVM_PATTERN_BITWISE_OP(operator^, prim::BitwiseXor);
+
+#undef TVM_PATTERN_BITWISE_OP
+
// logical expressions
TVM_PATTERN_BINARY_OP(operator>, prim::GT);
TVM_PATTERN_BINARY_OP(operator>=, prim::GE);
@@ -464,6 +483,37 @@ inline PNotExpr<TA> operator!(const Pattern<TA>& value) {
return PNotExpr<TA>(value.derived());
}
+/*!
+ * \brief Pattern bitwise not expression.
+ * \tparam TA The pattern type of the true operand.
+ */
+template <typename TA>
+class PBitwiseNotExpr : public Pattern<PBitwiseNotExpr<TA>> {
+ public:
+ explicit PBitwiseNotExpr(const TA& value) : value_(value) {}
+
+ void InitMatch_() const { value_.InitMatch_(); }
+
+ bool Match_(const ffi::ObjectRef& node) const {
+ if (const prim::BitwiseNotNode* ptr = node.as<prim::BitwiseNotNode>()) {
+ if (!value_.Match_(ptr->a)) return false;
+ return true;
+ } else {
+ return false;
+ }
+ }
+
+ PrimExpr Eval() const { return prim::BitwiseNot(value_.Eval()); }
+
+ private:
+ typename TA::Nested value_;
+};
+
+template <typename TA>
+inline PBitwiseNotExpr<TA> operator~(const Pattern<TA>& value) {
+ return PBitwiseNotExpr<TA>(value.derived());
+}
+
// select
/*!
* \brief Pattern select expression.
@@ -778,40 +828,6 @@ class PCallExpr : public Pattern<PCallExpr<Op, TArgs...>> {
std::tuple<typename TArgs::Nested...> args_;
};
-// arithemetic intrinsics
-#define TVM_PATTERN_BINARY_INTRIN(FuncName, OpName, IntrinOpName)
\
- struct OpName {
\
- static PrimExpr Eval(ffi::Array<PrimExpr> args) {
\
- return Call(args[0].ty(), GetOp(), args).as_or_throw<PrimExpr>();
\
- }
\
- static const Op& GetOp() { return prim::builtin::IntrinOpName(); }
\
- };
\
- template <typename TA, typename TB>
\
- inline PCallExpr<OpName, TA, TB> FuncName(const Pattern<TA>& a, const
Pattern<TB>& b) { \
- return PCallExpr<OpName, TA, TB>(a.derived(), b.derived());
\
- }
-
-TVM_PATTERN_BINARY_INTRIN(operator<<, PLeftShiftOp, shift_left);
-TVM_PATTERN_BINARY_INTRIN(operator>>, PRightShiftOp, shift_right);
-TVM_PATTERN_BINARY_INTRIN(operator&, PBitwiseAndOp, bitwise_and);
-TVM_PATTERN_BINARY_INTRIN(operator|, PBitwiseOrOp, bitwise_or);
-TVM_PATTERN_BINARY_INTRIN(operator^, PBitwiseXorOp, bitwise_xor);
-
-// unary intrinsics
-#define TVM_PATTERN_UNARY_INTRIN(FuncName, OpName, IntrinOpName) \
- struct OpName { \
- static PrimExpr Eval(ffi::Array<PrimExpr> args) { \
- return Call(args[0].ty(), GetOp(), args).as_or_throw<PrimExpr>(); \
- } \
- static const Op& GetOp() { return prim::builtin::IntrinOpName(); } \
- }; \
- template <typename TA> \
- inline PCallExpr<OpName, TA> FuncName(const Pattern<TA>& a) { \
- return PCallExpr<OpName, TA>(a.derived()); \
- }
-
-TVM_PATTERN_UNARY_INTRIN(operator~, PBitwiseNotOp, bitwise_not);
-
// if_then_else
struct PIfThenElseOp {
static PrimExpr Eval(ffi::Array<PrimExpr> args) {
diff --git a/src/sym/rewrite_simplify.cc b/src/sym/rewrite_simplify.cc
index a340654c9a..1f9491e2d9 100644
--- a/src/sym/rewrite_simplify.cc
+++ b/src/sym/rewrite_simplify.cc
@@ -2411,6 +2411,30 @@ UnchangedOr<PrimExpr>
RewriteSimplifier::Impl::Mutate_(const prim::SelectNode* o
return ret;
}
+UnchangedOr<PrimExpr> RewriteSimplifier::Impl::Mutate_(const prim::LShiftNode*
op,
+ InplaceMode
inplace_mode) {
+ PrimExpr ret =
+ SimplifierBase::Mutate_(op,
inplace_mode).ValueOrUnchanged(ffi::GetRef<PrimExpr>(op));
+ op = ret.as<prim::LShiftNode>();
+ if (op && op->a.as<IntImmNode>() && op->b.as<IntImmNode>()) {
+ // The operator overload eagerly folds constant operands.
+ return op->a << op->b;
+ }
+ return ret;
+}
+
+UnchangedOr<PrimExpr> RewriteSimplifier::Impl::Mutate_(const prim::RShiftNode*
op,
+ InplaceMode
inplace_mode) {
+ PrimExpr ret =
+ SimplifierBase::Mutate_(op,
inplace_mode).ValueOrUnchanged(ffi::GetRef<PrimExpr>(op));
+ op = ret.as<prim::RShiftNode>();
+ if (op && op->a.as<IntImmNode>() && op->b.as<IntImmNode>()) {
+ // The operator overload eagerly folds constant operands.
+ return op->a >> op->b;
+ }
+ return ret;
+}
+
UnchangedOr<Expr> RewriteSimplifier::Impl::Mutate_(const CallNode* op,
InplaceMode inplace_mode) {
// add condition context to if_then_else
Expr expr = SimplifierBase::Mutate_(op,
inplace_mode).ValueOrUnchanged(ffi::GetRef<Expr>(op));
@@ -2425,16 +2449,6 @@ UnchangedOr<Expr> RewriteSimplifier::Impl::Mutate_(const
CallNode* op, InplaceMo
if (op->op.same_as(prim::builtin::likely()) &&
prim::is_const_int(op->args[0].as_or_throw<PrimExpr>())) {
return op->args[0].as_or_throw<PrimExpr>();
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- if (op->args[0].as<IntImmNode>() && op->args[1].as<IntImmNode>()) {
- // the operator overload will eagerly constant fold.
- return op->args[0].as_or_throw<PrimExpr>() >>
op->args[1].as_or_throw<PrimExpr>();
- }
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- if (op->args[0].as<IntImmNode>() && op->args[1].as<IntImmNode>()) {
- // the operator overload will eagerly constant fold.
- return op->args[0].as_or_throw<PrimExpr>() <<
op->args[1].as_or_throw<PrimExpr>();
- }
}
static const Op& ceil_op = prim::builtin::ceil();
static const Op& log2_op = prim::builtin::log2();
diff --git a/src/sym/rewrite_simplify.h b/src/sym/rewrite_simplify.h
index c51c58b32b..456bb0dcf5 100644
--- a/src/sym/rewrite_simplify.h
+++ b/src/sym/rewrite_simplify.h
@@ -96,6 +96,8 @@ class RewriteSimplifier::Impl : public SimplifierBase {
UnchangedOr<Expr> Mutate_(const CallNode* op, InplaceMode inplace_mode)
override;
UnchangedOr<Expr> Mutate_(const VarNode* op, InplaceMode inplace_mode)
override;
UnchangedOr<PrimExpr> Mutate_(const prim::AddNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::LShiftNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::RShiftNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::SubNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::MulNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::DivNode* op, InplaceMode
inplace_mode) override;
diff --git a/src/sym/z3_prover.cc b/src/sym/z3_prover.cc
index d20eac26bb..8bf7450c6d 100644
--- a/src/sym/z3_prover.cc
+++ b/src/sym/z3_prover.cc
@@ -741,7 +741,11 @@ class Z3Prover::Impl : tvm::ExprFunctor<z3::expr(const
Expr&)> {
if (memo_.count(e)) {
return false;
}
- return e->IsInstance<CallNode>() || e->IsInstance<TensorLoadNode>() ||
+ // Keep the conservative shortcut used when bitwise operations were calls.
+ return e->IsInstance<CallNode>() || e->IsInstance<prim::LShiftNode>() ||
+ e->IsInstance<prim::RShiftNode>() ||
e->IsInstance<prim::BitwiseAndNode>() ||
+ e->IsInstance<prim::BitwiseOrNode>() ||
e->IsInstance<prim::BitwiseXorNode>() ||
+ e->IsInstance<prim::BitwiseNotNode>() ||
e->IsInstance<TensorLoadNode>() ||
(e->IsInstance<prim::CastNode>() &&
!IsZ3SupportedExpr(e.as_or_throw<prim::Cast>()->value.get()));
}
@@ -884,23 +888,25 @@ class Z3Prover::Impl : tvm::ExprFunctor<z3::expr(const
Expr&)> {
}
z3::expr Dispatch_(const IntImmNode* op) override { return
IntegerValue(*ctx, op->value); }
- // Bitwise operations
+ z3::expr Dispatch_(const prim::BitwiseAndNode* op) override {
+ return VisitBitwiseOp(z3::operator&, op, op->a, op->b);
+ }
+ z3::expr Dispatch_(const prim::BitwiseOrNode* op) override {
+ return VisitBitwiseOp(z3::operator|, op, op->a, op->b);
+ }
+ z3::expr Dispatch_(const prim::BitwiseXorNode* op) override {
+ return VisitBitwiseOp(z3::operator^, op, op->a, op->b);
+ }
+ z3::expr Dispatch_(const prim::LShiftNode* op) override {
+ return VisitShiftOp(z3::shl, op, op->a, op->b);
+ }
+ z3::expr Dispatch_(const prim::RShiftNode* op) override {
+ return VisitShiftOp(z3::ashr, op, op->a, op->b);
+ }
+
z3::expr Dispatch_(const CallNode* op) override {
- // Check if this is a bitwise operation
- if (op->op.same_as(prim::builtin::bitwise_and())) {
- return VisitBitwiseOp(z3::operator&, op);
- } else if (op->op.same_as(prim::builtin::bitwise_or())) {
- return VisitBitwiseOp(z3::operator|, op);
- } else if (op->op.same_as(prim::builtin::bitwise_xor())) {
- return VisitBitwiseOp(z3::operator^, op);
- } else if (op->op.same_as(prim::builtin::bitwise_not())) {
- return VisitBitwiseNotOp(op);
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- return VisitShiftOp(z3::shl, op);
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- return VisitShiftOp(z3::ashr, op);
- } else if (op->op.same_as(prim::builtin::if_then_else()) &&
op->args.size() == 3 &&
- IsZ3SupportedExpr(op->args[1].get()) &&
IsZ3SupportedExpr(op->args[2].get())) {
+ if (op->op.same_as(prim::builtin::if_then_else()) && op->args.size() == 3
&&
+ IsZ3SupportedExpr(op->args[1].get()) &&
IsZ3SupportedExpr(op->args[2].get())) {
// tir.if_then_else(cond, a, b) is a select-like ternary.
return z3::ite(VisitBool(op->args[0].as_or_throw<PrimExpr>()),
VisitInt(op->args[1].as_or_throw<PrimExpr>()),
@@ -912,15 +918,8 @@ class Z3Prover::Impl : tvm::ExprFunctor<z3::expr(const
Expr&)> {
}
/// @brief Helper function to visit binary bitwise operations
- z3::expr VisitBitwiseOp(z3::expr (*op_func)(const z3::expr&, const
z3::expr&),
- const CallNode* op) {
- if (op->args.size() != 2) {
- LOG(FATAL) << "Binary bitwise operation expects 2 arguments, got " <<
op->args.size();
- TVM_FFI_UNREACHABLE();
- }
-
- PrimExpr a = op->args[0].as_or_throw<PrimExpr>();
- PrimExpr b = op->args[1].as_or_throw<PrimExpr>();
+ z3::expr VisitBitwiseOp(z3::expr (*op_func)(const z3::expr&, const
z3::expr&), const ExprNode* op,
+ const PrimExpr& a, const PrimExpr& b) {
unsigned bit_width = std::max(a.ty().bits(), b.ty().bits());
if (IsZ3SupportedExpr(a.get()) && IsZ3SupportedExpr(b.get())) {
@@ -932,13 +931,8 @@ class Z3Prover::Impl : tvm::ExprFunctor<z3::expr(const
Expr&)> {
}
/// @brief Helper function to visit unary bitwise not operation
- z3::expr VisitBitwiseNotOp(const CallNode* op) {
- if (op->args.size() != 1) {
- LOG(FATAL) << "Bitwise not operation expects 1 argument, got " <<
op->args.size();
- TVM_FFI_UNREACHABLE();
- }
-
- PrimExpr a = op->args[0].as_or_throw<PrimExpr>();
+ z3::expr Dispatch_(const prim::BitwiseNotNode* op) override {
+ const PrimExpr& a = op->a;
if (IsZ3SupportedExpr(a.get())) {
// Cast integer to bit-vector, apply bitwise not, then cast back.
@@ -952,15 +946,8 @@ class Z3Prover::Impl : tvm::ExprFunctor<z3::expr(const
Expr&)> {
}
/// @brief Helper function to visit shift operations
- z3::expr VisitShiftOp(z3::expr (*op_func)(const z3::expr&, const z3::expr&),
const CallNode* op) {
- if (op->args.size() != 2) {
- LOG(FATAL) << "Shift operation expects 2 arguments, got " <<
op->args.size();
- TVM_FFI_UNREACHABLE();
- }
-
- PrimExpr a = op->args[0].as_or_throw<PrimExpr>();
- PrimExpr b = op->args[1].as_or_throw<PrimExpr>();
-
+ z3::expr VisitShiftOp(z3::expr (*op_func)(const z3::expr&, const z3::expr&),
const ExprNode* op,
+ const PrimExpr& a, const PrimExpr& b) {
// Shift operations require integer types for both operands
if (IsZ3SupportedExpr(a.get()) && IsZ3SupportedExpr(b.get())) {
z3::expr a_expr = VisitInt(a);
diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc
index 575388e0a4..ed727151d9 100644
--- a/src/target/llvm/codegen_llvm.cc
+++ b/src/target/llvm/codegen_llvm.cc
@@ -1398,22 +1398,6 @@ llvm::Value* CodeGenLLVM::CreateIntrinsic(const
CallNode* op) {
}
}
return builder_->CreateCall(f, arg_value);
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- return builder_->CreateAnd(MakeValue(args[0]), MakeValue(args[1]));
- } else if (op->op.same_as(prim::builtin::bitwise_or())) {
- return builder_->CreateOr(MakeValue(args[0]), MakeValue(args[1]));
- } else if (op->op.same_as(prim::builtin::bitwise_not())) {
- return builder_->CreateNot(MakeValue(args[0]));
- } else if (op->op.same_as(prim::builtin::bitwise_xor())) {
- return builder_->CreateXor(MakeValue(args[0]), MakeValue(args[1]));
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- return builder_->CreateShl(MakeValue(args[0]), MakeValue(args[1]));
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- if
(args[0].as_or_throw<PrimExpr>().ty().MatchesCode(DLDataTypeCode::kDLInt)) {
- return builder_->CreateAShr(MakeValue(args[0]), MakeValue(args[1]));
- } else {
- return builder_->CreateLShr(MakeValue(args[0]), MakeValue(args[1]));
- }
} else if (op->op.same_as(tirx::builtin::tvm_storage_sync())) {
return CreateStorageSync(op);
} else if (op->op.same_as(tirx::builtin::address_of())) {
@@ -1725,6 +1709,34 @@ llvm::Value* CodeGenLLVM::Dispatch_(const prim::NotNode*
op) {
return builder_->CreateNot(MakeValue(op->a));
}
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::LShiftNode* op) {
+ return builder_->CreateShl(MakeValue(op->a), MakeValue(op->b));
+}
+
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::BitwiseAndNode* op) {
+ return builder_->CreateAnd(MakeValue(op->a), MakeValue(op->b));
+}
+
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::BitwiseOrNode* op) {
+ return builder_->CreateOr(MakeValue(op->a), MakeValue(op->b));
+}
+
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::BitwiseXorNode* op) {
+ return builder_->CreateXor(MakeValue(op->a), MakeValue(op->b));
+}
+
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::RShiftNode* op) {
+ if (op->a.ty().MatchesCode(DLDataTypeCode::kDLInt)) {
+ return builder_->CreateAShr(MakeValue(op->a), MakeValue(op->b));
+ } else {
+ return builder_->CreateLShr(MakeValue(op->a), MakeValue(op->b));
+ }
+}
+
+llvm::Value* CodeGenLLVM::Dispatch_(const prim::BitwiseNotNode* op) {
+ return builder_->CreateNot(MakeValue(op->a));
+}
+
llvm::Value* CodeGenLLVM::Dispatch_(const prim::SelectNode* op) {
return builder_->CreateSelect(MakeValue(op->condition),
MakeValue(op->true_value),
MakeValue(op->false_value));
diff --git a/src/target/llvm/codegen_llvm.h b/src/target/llvm/codegen_llvm.h
index 0e2d11e10c..98f991b1ba 100644
--- a/src/target/llvm/codegen_llvm.h
+++ b/src/target/llvm/codegen_llvm.h
@@ -223,6 +223,12 @@ class CodeGenLLVM : public
tirx::ExprFunctor<llvm::Value*(const Expr&)>,
llvm::Value* Dispatch_(const prim::AndNode* op) override;
llvm::Value* Dispatch_(const prim::OrNode* op) override;
llvm::Value* Dispatch_(const prim::NotNode* op) override;
+ llvm::Value* Dispatch_(const prim::LShiftNode* op) override;
+ llvm::Value* Dispatch_(const prim::RShiftNode* op) override;
+ llvm::Value* Dispatch_(const prim::BitwiseAndNode* op) override;
+ llvm::Value* Dispatch_(const prim::BitwiseOrNode* op) override;
+ llvm::Value* Dispatch_(const prim::BitwiseXorNode* op) override;
+ llvm::Value* Dispatch_(const prim::BitwiseNotNode* op) override;
llvm::Value* Dispatch_(const prim::SelectNode* op) override;
llvm::Value* Dispatch_(const prim::LetNode* op) override;
llvm::Value* Dispatch_(const TensorLoadNode* op) override;
diff --git a/src/target/source/codegen_c.cc b/src/target/source/codegen_c.cc
index f2cea3de5c..a496cee1b2 100644
--- a/src/target/source/codegen_c.cc
+++ b/src/target/source/codegen_c.cc
@@ -577,22 +577,6 @@ inline void PrintBinaryExpr(const T* op, const char* opstr,
}
}
-inline void PrintBinaryIntrinsic(const CallNode* op, const char* opstr,
- std::ostream& os, // NOLINT(*)
- CodeGenC* p) {
- PrimType op_ty = op->ty.as_or_throw<PrimType>();
- if (op_ty.lanes() == 1) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 2U);
- os << '(';
- p->PrintExpr(op->args[0], os);
- os << opstr;
- p->PrintExpr(op->args[1], os);
- os << ')';
- } else {
- p->PrintVecBinaryOp(opstr, op_ty, op->args[0].as_or_throw<PrimExpr>(),
- op->args[1].as_or_throw<PrimExpr>(), os);
- }
-}
void CodeGenC::Dispatch_(const prim::CastNode* op, std::ostream& os) { //
NOLINT(*)
std::stringstream value;
this->PrintExpr(op->value, value);
@@ -667,6 +651,27 @@ void CodeGenC::Dispatch_(const prim::NotNode* op,
std::ostream& os) { // NOLINT
PrintExpr(op->a, os);
}
+void CodeGenC::Dispatch_(const prim::LShiftNode* op, std::ostream& os) { //
NOLINT(*)
+ PrintBinaryExpr(op, "<<", os, this);
+}
+void CodeGenC::Dispatch_(const prim::RShiftNode* op, std::ostream& os) { //
NOLINT(*)
+ PrintBinaryExpr(op, ">>", os, this);
+}
+void CodeGenC::Dispatch_(const prim::BitwiseAndNode* op, std::ostream& os) {
// NOLINT(*)
+ PrintBinaryExpr(op, "&", os, this);
+}
+void CodeGenC::Dispatch_(const prim::BitwiseOrNode* op, std::ostream& os) {
// NOLINT(*)
+ PrintBinaryExpr(op, "|", os, this);
+}
+void CodeGenC::Dispatch_(const prim::BitwiseXorNode* op, std::ostream& os) {
// NOLINT(*)
+ PrintBinaryExpr(op, "^", os, this);
+}
+void CodeGenC::Dispatch_(const prim::BitwiseNotNode* op, std::ostream& os) {
// NOLINT(*)
+ os << "(~";
+ PrintExpr(op->a, os);
+ os << ')';
+}
+
void CodeGenC::PrintCallExtern(Type ret_type, ffi::String global_symbol,
const ffi::Array<Expr>& args, bool
skip_first_arg,
std::ostream& os) { // NOLINT(*)
@@ -733,21 +738,6 @@ void CodeGenC::Dispatch_(const CallNode* op, std::ostream&
os) { // NOLINT(*)
// call extern if the op itself have a global symbol.
ffi::Array<Expr> args = op->args;
this->PrintCallExtern(op->ty, op_attr_global_symbol_[call_op], args,
false, os);
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- PrintBinaryIntrinsic(op, " & ", os, this);
- } else if (op->op.same_as(prim::builtin::bitwise_xor())) {
- PrintBinaryIntrinsic(op, " ^ ", os, this);
- } else if (op->op.same_as(prim::builtin::bitwise_or())) {
- PrintBinaryIntrinsic(op, " | ", os, this);
- } else if (op->op.same_as(prim::builtin::bitwise_not())) {
- TVM_FFI_ICHECK_EQ(op->args.size(), 1U);
- os << "(~";
- this->PrintExpr(op->args[0], os);
- os << ')';
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- PrintBinaryIntrinsic(op, " << ", os, this);
- } else if (op->op.same_as(prim::builtin::shift_right())) {
- PrintBinaryIntrinsic(op, " >> ", os, this);
} else if (op->op.same_as(prim::builtin::if_then_else())) {
// conditional that skips eval if cond evals to false
std::string result = name_supply_->FreshName("condval");
diff --git a/src/target/source/codegen_c.h b/src/target/source/codegen_c.h
index dbe9dbc28a..f6707438b7 100644
--- a/src/target/source/codegen_c.h
+++ b/src/target/source/codegen_c.h
@@ -168,34 +168,40 @@ class CodeGenC : public tirx::ExprFunctor<void(const
Expr&, std::ostream&)>,
*/
virtual void InitFuncState(const PrimFunc& f);
// expression
- void Dispatch_(const VarNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const TensorLoadNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::LetNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const CallNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::AddNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::SubNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::MulNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::DivNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::ModNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::MinNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::MaxNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::EQNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::NENode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::LTNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::LENode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::GTNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::GENode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::AndNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::OrNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::CastNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::NotNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::SelectNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::RampNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::ShuffleNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const prim::BroadcastNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const IntImmNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const FloatImmNode* op, std::ostream& os) override;
// NOLINT(*)
- void Dispatch_(const StringImmNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const VarNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const TensorLoadNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::LetNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const CallNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::AddNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::SubNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::MulNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::DivNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::ModNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::MinNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::MaxNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::EQNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::NENode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::LTNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::LENode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::GTNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::GENode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::AndNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::OrNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::CastNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::NotNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::LShiftNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::RShiftNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::BitwiseAndNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::BitwiseOrNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::BitwiseXorNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::BitwiseNotNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::SelectNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::RampNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::ShuffleNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const prim::BroadcastNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const IntImmNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const FloatImmNode* op, std::ostream& os) override;
// NOLINT(*)
+ void Dispatch_(const StringImmNode* op, std::ostream& os) override;
// NOLINT(*)
// statment
void Dispatch_(const BindNode* op) override;
void Dispatch_(const BufferStoreNode* op) override;
diff --git a/src/tirx/analysis/filter_canonical.cc
b/src/tirx/analysis/filter_canonical.cc
index eeeda18780..9f0b570dbc 100644
--- a/src/tirx/analysis/filter_canonical.cc
+++ b/src/tirx/analysis/filter_canonical.cc
@@ -37,14 +37,6 @@ namespace tirx {
namespace {
-// Recognized conjunction shapes: logical-And and bitwise-And calls.
-// Mirrors FlattenConjuncts in tile_primitive_dispatch.cc so the classifier
-// accepts the same set of "fully conjunctive" predicates that the existing
-// pass-internal helpers do.
-bool IsBitwiseAndCall(const CallNode* call) {
- return call->op.same_as(prim::builtin::bitwise_and()) && call->args.size()
== 2;
-}
-
bool IsPtxElectSyncCall(const CallNode* call) {
static const Op& ptx_elect_sync_op = Op::Get("tirx.cuda.elect_sync");
return call->op.same_as(ptx_elect_sync_op);
@@ -63,6 +55,10 @@ PrimExpr StripCast(const PrimExpr& expr) {
return cur;
}
+// Recognized conjunction shapes: logical-And and bitwise-And nodes.
+// Mirrors FlattenConjuncts in tile_primitive_dispatch.cc so the classifier
+// accepts the same set of "fully conjunctive" predicates that the existing
+// pass-internal helpers do.
void FlattenConjuncts(const PrimExpr& pred, std::vector<PrimExpr>* out) {
PrimExpr stripped = StripCast(pred);
if (const auto* and_node = stripped.as<prim::AndNode>()) {
@@ -70,12 +66,10 @@ void FlattenConjuncts(const PrimExpr& pred,
std::vector<PrimExpr>* out) {
FlattenConjuncts(and_node->b, out);
return;
}
- if (const auto* call = stripped.as<CallNode>()) {
- if (IsBitwiseAndCall(call)) {
- FlattenConjuncts(call->args[0].as_or_throw<PrimExpr>(), out);
- FlattenConjuncts(call->args[1].as_or_throw<PrimExpr>(), out);
- return;
- }
+ if (const auto* and_node = stripped.as<prim::BitwiseAndNode>()) {
+ FlattenConjuncts(and_node->a, out);
+ FlattenConjuncts(and_node->b, out);
+ return;
}
out->push_back(stripped);
}
diff --git a/src/tirx/analysis/filter_canonical.h
b/src/tirx/analysis/filter_canonical.h
index 23fbdd7c4f..9e35ca2252 100644
--- a/src/tirx/analysis/filter_canonical.h
+++ b/src/tirx/analysis/filter_canonical.h
@@ -134,7 +134,7 @@ using ScopeIdPredicate = std::function<bool(const Var&)>;
*
* Implementation notes:
* - Conjunction is recognized via both `tir::And` nodes and
- * `tirx.bitwise_and` calls (matching existing FlattenConjuncts behavior
+ * `prim.BitwiseAnd` nodes (matching existing FlattenConjuncts behavior
* in tile_primitive_dispatch.cc).
* - Comparison atoms with `const <op> var` are mirrored so the
* `scopeid_var` is on the LHS of the returned atom.
diff --git a/src/tirx/ir/data_type_rewriter.cc
b/src/tirx/ir/data_type_rewriter.cc
index b3f604fd57..2de7d9c75a 100644
--- a/src/tirx/ir/data_type_rewriter.cc
+++ b/src/tirx/ir/data_type_rewriter.cc
@@ -252,11 +252,60 @@ TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::LTNode,
operator<); // NOLINT
TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::GTNode, operator>); //
NOLINT(*)
TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::GENode, operator>=);
+TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::BitwiseAndNode, bitwise_and);
+TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::BitwiseOrNode, bitwise_or);
+TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH(prim::BitwiseXorNode, bitwise_xor);
+
#undef TVM_DEFINE_BIOP_EXPR_MUTATE_WITH_TYPE_MATCH
+UnchangedOr<PrimExpr> DataTypeLegalizer::Mutate_(const prim::LShiftNode* op,
+ InplaceMode inplace_mode) {
+ PrimType before_dtype = op->a.ty();
+ // Preserve the original operands while computing the narrowed shift.
+ PrimExpr lhs = Mutate(op->a, InplaceMode::kDisallow).ValueOrUnchanged(op->a);
+ PrimExpr rhs = Mutate(op->b, InplaceMode::kDisallow).ValueOrUnchanged(op->b);
+ PrimType after_dtype = lhs.ty();
+ if (ShouldClampShiftAmounts() && before_dtype.code() ==
DLDataTypeCode::kDLInt &&
+ after_dtype.code() == DLDataTypeCode::kDLInt && before_dtype.bits() >
after_dtype.bits()) {
+ // Values fit in the narrowed dtype. Clamp lane-wise to keep dynamic and
+ // vector shift amounts below its width, preserving representable results.
+ rhs = min(rhs, MakeConst(rhs.ty(), after_dtype.bits() - 1, op->span),
op->span);
+ }
+ if (lhs.same_as(op->a) && rhs.same_as(op->b) && lhs.ty() == rhs.ty()) {
+ return ffi::Unchanged();
+ }
+ return left_shift(lhs, rhs, op->span);
+}
+
+UnchangedOr<PrimExpr> DataTypeLegalizer::Mutate_(const prim::RShiftNode* op,
+ InplaceMode inplace_mode) {
+ PrimType before_dtype = op->a.ty();
+ // Preserve the original operands while computing the narrowed shift.
+ PrimExpr lhs = Mutate(op->a, InplaceMode::kDisallow).ValueOrUnchanged(op->a);
+ PrimExpr rhs = Mutate(op->b, InplaceMode::kDisallow).ValueOrUnchanged(op->b);
+ PrimType after_dtype = lhs.ty();
+ if (ShouldClampShiftAmounts() && before_dtype.code() ==
DLDataTypeCode::kDLInt &&
+ after_dtype.code() == DLDataTypeCode::kDLInt && before_dtype.bits() >
after_dtype.bits()) {
+ // Values fit in the narrowed dtype. Clamp lane-wise to keep dynamic and
+ // vector shift amounts below its width, preserving representable results.
+ rhs = min(rhs, MakeConst(rhs.ty(), after_dtype.bits() - 1, op->span),
op->span);
+ }
+ if (lhs.same_as(op->a) && rhs.same_as(op->b) && lhs.ty() == rhs.ty()) {
+ return ffi::Unchanged();
+ }
+ return right_shift(lhs, rhs, op->span);
+}
+
+UnchangedOr<PrimExpr> DataTypeLegalizer::Mutate_(const prim::BitwiseNotNode*
op,
+ InplaceMode inplace_mode) {
+ auto a = Mutate(op->a, inplace_mode);
+ if (a.UnchangedOrSameAs(op->a)) return ffi::Unchanged();
+ return prim::BitwiseNot(std::move(a).ValueOrUnchanged(op->a), op->span);
+}
+
UnchangedOr<Expr> DataTypeLegalizer::Mutate_(const CallNode* op, InplaceMode
inplace_mode) {
Call before = ffi::GetRef<Call>(op);
- // Keep the original argument dtypes available for shift and clz correction
below.
+ // Keep the original argument dtype available for clz correction below.
Expr e =
StmtExprMutator::Mutate_(op,
InplaceMode::kDisallow).ValueOrUnchanged(ffi::GetRef<Expr>(op));
op = e.as<CallNode>();
@@ -266,40 +315,6 @@ UnchangedOr<Expr> DataTypeLegalizer::Mutate_(const
CallNode* op, InplaceMode inp
return e;
}
PrimExpr prim_e = e.as_or_throw<PrimExpr>();
- if (op->op.same_as(prim::builtin::shift_right())) {
- PrimExpr lhs = op->args[0].as_or_throw<PrimExpr>();
- PrimExpr rhs = op->args[1].as_or_throw<PrimExpr>();
- PrimType before_dtype = before->args[0].as_or_throw<PrimExpr>().ty();
- PrimType after_dtype = lhs.ty();
- if (ShouldClampShiftAmounts() && before_dtype.code() ==
DLDataTypeCode::kDLInt &&
- after_dtype.code() == DLDataTypeCode::kDLInt && before_dtype.bits() >
after_dtype.bits()) {
- // Values are assumed to fit in the narrowed dtype. An arithmetic right
- // shift at or beyond its sign bit therefore has the same value as a
shift
- // by the new sign-bit position. Clamp lane-wise so dynamic and vector
- // shift amounts remain valid for the narrowed dtype.
- rhs = min(rhs, MakeConst(rhs.ty(), after_dtype.bits() - 1, op->span),
op->span);
- }
- return lhs >> rhs;
- } else if (op->op.same_as(prim::builtin::shift_left())) {
- PrimExpr lhs = op->args[0].as_or_throw<PrimExpr>();
- PrimExpr rhs = op->args[1].as_or_throw<PrimExpr>();
- PrimType before_dtype = before->args[0].as_or_throw<PrimExpr>().ty();
- PrimType after_dtype = lhs.ty();
- if (ShouldClampShiftAmounts() && before_dtype.code() ==
DLDataTypeCode::kDLInt &&
- after_dtype.code() == DLDataTypeCode::kDLInt && before_dtype.bits() >
after_dtype.bits()) {
- // Keep dynamic and vector shift amounts valid for the narrowed dtype.
Under the pass's
- // representability precondition, a left shift at or beyond the narrowed
width can only
- // produce a representable result when lhs is zero, so clamping does not
alter valid cases.
- rhs = min(rhs, MakeConst(rhs.ty(), after_dtype.bits() - 1, op->span),
op->span);
- }
- return lhs << rhs;
- } else if (op->op.same_as(prim::builtin::bitwise_and())) {
- return op->args[0].as_or_throw<PrimExpr>() &
op->args[1].as_or_throw<PrimExpr>();
- } else if (op->op.same_as(prim::builtin::bitwise_or())) {
- return op->args[0].as_or_throw<PrimExpr>() |
op->args[1].as_or_throw<PrimExpr>();
- } else if (op->op.same_as(prim::builtin::bitwise_xor())) {
- return op->args[0].as_or_throw<PrimExpr>() ^
op->args[1].as_or_throw<PrimExpr>();
- }
static const Op& pow_op = Op::Get("tirx.pow");
static const Op& clz_op = prim::builtin::clz();
if (op->op.same_as(pow_op)) {
diff --git a/src/tirx/ir/data_type_rewriter.h b/src/tirx/ir/data_type_rewriter.h
index a95873dd15..2150dac228 100644
--- a/src/tirx/ir/data_type_rewriter.h
+++ b/src/tirx/ir/data_type_rewriter.h
@@ -61,6 +61,12 @@ class DataTypeLegalizer : public StmtExprMutator {
UnchangedOr<PrimExpr> Mutate_(const prim::RampNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::BroadcastNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::ShuffleNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::LShiftNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::RShiftNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseAndNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseOrNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseXorNode* op, InplaceMode
inplace_mode) override;
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseNotNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::AddNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::SubNode* op, InplaceMode
inplace_mode) override;
UnchangedOr<PrimExpr> Mutate_(const prim::MulNode* op, InplaceMode
inplace_mode) override;
diff --git a/src/tirx/ir/specialize.cc b/src/tirx/ir/specialize.cc
index 55195e3dc6..66a06863b1 100644
--- a/src/tirx/ir/specialize.cc
+++ b/src/tirx/ir/specialize.cc
@@ -64,7 +64,7 @@ inline bool IsParam(const PrimFunc& func, const Var& param) {
if (a_unchanged && b_unchanged) {
\
return ffi::Unchanged();
\
} else {
\
- return BinaryFunc(a, b);
\
+ return BinaryFunc(a, b, op->span);
\
}
\
}
#define DEFINE_SPECIALIZER_UNARY_OP_MUTATE(UnaryNode, UnaryFunc)
\
@@ -75,7 +75,7 @@ inline bool IsParam(const PrimFunc& func, const Var& param) {
if (a_unchanged) {
\
return ffi::Unchanged();
\
} else {
\
- return UnaryFunc(a);
\
+ return UnaryFunc(a, op->span);
\
}
\
}
@@ -244,6 +244,12 @@ class PrimFuncSpecializer : public StmtExprMutator {
DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::AndNode, logical_and);
DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::OrNode, logical_or);
DEFINE_SPECIALIZER_UNARY_OP_MUTATE(prim::NotNode, logical_not);
+ DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::LShiftNode, left_shift);
+ DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::RShiftNode, right_shift);
+ DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::BitwiseAndNode, bitwise_and);
+ DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::BitwiseOrNode, bitwise_or);
+ DEFINE_SPECIALIZER_BINARY_OP_MUTATE(prim::BitwiseXorNode, bitwise_xor);
+ DEFINE_SPECIALIZER_UNARY_OP_MUTATE(prim::BitwiseNotNode, prim::BitwiseNot);
BufferVar MutateBuffer(const BufferVar& buffer) {
ffi::Any mapped = VarRemapGet(buffer);
if (mapped.type_index() != ffi::TypeIndex::kTVMFFINone) {
diff --git a/src/tirx/ir/tir_visitor_with_path.cc
b/src/tirx/ir/tir_visitor_with_path.cc
index 6cc64a5b87..6757e1b01a 100644
--- a/src/tirx/ir/tir_visitor_with_path.cc
+++ b/src/tirx/ir/tir_visitor_with_path.cc
@@ -353,6 +353,12 @@ DEFINE_BINOP_VISIT_(prim::GENode);
DEFINE_BINOP_VISIT_(prim::AndNode);
DEFINE_BINOP_VISIT_(prim::OrNode);
+DEFINE_BINOP_VISIT_(prim::LShiftNode);
+DEFINE_BINOP_VISIT_(prim::RShiftNode);
+DEFINE_BINOP_VISIT_(prim::BitwiseAndNode);
+DEFINE_BINOP_VISIT_(prim::BitwiseOrNode);
+DEFINE_BINOP_VISIT_(prim::BitwiseXorNode);
+
#undef DEFINE_BINOP_VISIT_
void TIRVisitorWithPath::Dispatch_(const IntImmNode* op, AccessPath path) {}
@@ -363,6 +369,10 @@ void TIRVisitorWithPath::Dispatch_(const prim::CastNode*
op, AccessPath path) {
Visit(op->value, path->Attr("value"));
}
+void TIRVisitorWithPath::Dispatch_(const prim::BitwiseNotNode* op, AccessPath
path) {
+ Visit(op->a, path->Attr("a"));
+}
+
void TIRVisitorWithPath::Dispatch_(const prim::NotNode* op, AccessPath path) {
Visit(op->a, path->Attr("a"));
}
diff --git a/src/tirx/ir/tir_visitor_with_path.h
b/src/tirx/ir/tir_visitor_with_path.h
index 3bef3ad4bd..56305d8698 100644
--- a/src/tirx/ir/tir_visitor_with_path.h
+++ b/src/tirx/ir/tir_visitor_with_path.h
@@ -186,6 +186,12 @@ class TIRVisitorWithPath : protected
ExprFunctor<void(const Expr&, ffi::reflecti
void Dispatch_(const prim::AndNode* op, ffi::reflection::AccessPath path)
override;
void Dispatch_(const prim::OrNode* op, ffi::reflection::AccessPath path)
override;
void Dispatch_(const prim::CastNode* op, ffi::reflection::AccessPath path)
override;
+ void Dispatch_(const prim::LShiftNode* op, ffi::reflection::AccessPath path)
override;
+ void Dispatch_(const prim::RShiftNode* op, ffi::reflection::AccessPath path)
override;
+ void Dispatch_(const prim::BitwiseAndNode* op, ffi::reflection::AccessPath
path) override;
+ void Dispatch_(const prim::BitwiseOrNode* op, ffi::reflection::AccessPath
path) override;
+ void Dispatch_(const prim::BitwiseXorNode* op, ffi::reflection::AccessPath
path) override;
+ void Dispatch_(const prim::BitwiseNotNode* op, ffi::reflection::AccessPath
path) override;
void Dispatch_(const prim::NotNode* op, ffi::reflection::AccessPath path)
override;
void Dispatch_(const prim::SelectNode* op, ffi::reflection::AccessPath path)
override;
void Dispatch_(const prim::RampNode* op, ffi::reflection::AccessPath path)
override;
diff --git a/src/tirx/op/builtin.cc b/src/tirx/op/builtin.cc
index d47f9f75ce..84db3c9749 100644
--- a/src/tirx/op/builtin.cc
+++ b/src/tirx/op/builtin.cc
@@ -40,12 +40,6 @@ TVM_FFI_STATIC_INIT_BLOCK() {
.set_attr<TIRxOpCategory>("TIRxOpCategory", ffi::String("builtin"), 1)
PRIM_SCRIPT_BUILTIN(likely);
- PRIM_SCRIPT_BUILTIN(bitwise_and);
- PRIM_SCRIPT_BUILTIN(bitwise_or);
- PRIM_SCRIPT_BUILTIN(bitwise_xor);
- PRIM_SCRIPT_BUILTIN(bitwise_not);
- PRIM_SCRIPT_BUILTIN(shift_left);
- PRIM_SCRIPT_BUILTIN(shift_right);
PRIM_SCRIPT_BUILTIN(if_then_else);
PRIM_SCRIPT_BUILTIN(vscale);
PRIM_SCRIPT_BUILTIN(ceil);
diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc
index 5a1ba7885a..e3a77b19ab 100644
--- a/src/tirx/script/printer/expr.cc
+++ b/src/tirx/script/printer/expr.cc
@@ -151,6 +151,17 @@ TVM_FFI_STATIC_INIT_BLOCK() {
});
}
+TVM_FFI_STATIC_INIT_BLOCK() {
+ IRDocsifier::vtable().set_dispatch<prim::BitwiseNot>(
+ "", [](prim::BitwiseNot node, AccessPath p, IRDocsifier d) -> Doc {
+ ExprDoc a = d->AsDoc<ExprDoc>(node->a, p->Attr("a"));
+ if (a->IsInstance<LiteralDocNode>()) {
+ return TIR(d, "BitwiseNot")->Call({a});
+ }
+ return OperationDoc(OperationDocNode::Kind::kInvert, {a});
+ });
+}
+
TVM_FFI_STATIC_INIT_BLOCK() {
IRDocsifier::vtable().set_dispatch<prim::Not>(
"", [](prim::Not node, AccessPath p, IRDocsifier d) -> Doc {
@@ -529,6 +540,15 @@ TVM_FFI_STATIC_INIT_BLOCK() {
kFloorDiv);
TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(FloorMod, prim::FloorModNode,
floormod, "FloorMod",
kMod);
+ TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(LShift, prim::LShiftNode,
left_shift, "LShift", kLShift);
+ TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(RShift, prim::RShiftNode,
right_shift, "RShift",
+ kRShift);
+ TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(BitwiseAnd, prim::BitwiseAndNode,
bitwise_and,
+ "BitwiseAnd", kBitAnd);
+ TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(BitwiseOr, prim::BitwiseOrNode,
bitwise_or, "BitwiseOr",
+ kBitOr);
+ TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(BitwiseXor, prim::BitwiseXorNode,
bitwise_xor,
+ "BitwiseXor", kBitXor);
TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(LT, prim::LTNode, less, "LT", kLt);
TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(LE, prim::LENode, less_equal, "LE",
kLtE);
TVM_SCRIPT_PRINTER_DEF_BINARY_WITH_SUGAR(EQ, prim::EQNode, equal, "EQ", kEq);
@@ -556,6 +576,12 @@ TVM_SCRIPT_REPR(prim::DivNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::ModNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::FloorDivNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::FloorModNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::LShiftNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::RShiftNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::BitwiseAndNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::BitwiseOrNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::BitwiseXorNode, ReprPrintTIR);
+TVM_SCRIPT_REPR(prim::BitwiseNotNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::MinNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::MaxNode, ReprPrintTIR);
TVM_SCRIPT_REPR(prim::LTNode, ReprPrintTIR);
diff --git a/src/tirx/transform/common_subexpr_elim.cc
b/src/tirx/transform/common_subexpr_elim.cc
index 9a6953efde..420a1f2141 100644
--- a/src/tirx/transform/common_subexpr_elim.cc
+++ b/src/tirx/transform/common_subexpr_elim.cc
@@ -499,8 +499,18 @@ class CSEPlanner : public StmtExprVisitor {
CSE_VISIT_BINARY(prim::GENode)
CSE_VISIT_BINARY(prim::AndNode)
CSE_VISIT_BINARY(prim::OrNode)
+ CSE_VISIT_BINARY(prim::LShiftNode)
+ CSE_VISIT_BINARY(prim::RShiftNode)
+ CSE_VISIT_BINARY(prim::BitwiseAndNode)
+ CSE_VISIT_BINARY(prim::BitwiseOrNode)
+ CSE_VISIT_BINARY(prim::BitwiseXorNode)
#undef CSE_VISIT_BINARY
+ ffi::Optional<VisitInterrupt> Visit_(const prim::BitwiseNotNode* op)
override {
+ TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op));
+ RecordExpr(ffi::GetRef<PrimExpr>(op), {op->a});
+ return std::nullopt;
+ }
ffi::Optional<VisitInterrupt> Visit_(const prim::NotNode* op) override {
TVM_FFI_S_VISIT_MAYBE_EARLY_RETURN(StmtExprVisitor::Visit_(op));
RecordExpr(ffi::GetRef<PrimExpr>(op), {op->a});
diff --git a/src/tirx/transform/tile_primitive_dispatch.cc
b/src/tirx/transform/tile_primitive_dispatch.cc
index 9de41412a3..7ffbd25ed4 100644
--- a/src/tirx/transform/tile_primitive_dispatch.cc
+++ b/src/tirx/transform/tile_primitive_dispatch.cc
@@ -1227,22 +1227,16 @@ class TilePrimitiveDispatcher : public StmtExprMutator {
return false;
}
- static bool IsBitwiseAndCall(const CallNode* call) {
- return call->op.same_as(prim::builtin::bitwise_and()) && call->args.size()
== 2;
- }
-
void FlattenConjuncts(const PrimExpr& pred, std::vector<PrimExpr>* out)
const {
if (const auto* and_node = pred.as<prim::AndNode>()) {
FlattenConjuncts(and_node->a, out);
FlattenConjuncts(and_node->b, out);
return;
}
- if (const auto* call = pred.as<CallNode>()) {
- if (IsBitwiseAndCall(call)) {
- FlattenConjuncts(call->args[0].as_or_throw<PrimExpr>(), out);
- FlattenConjuncts(call->args[1].as_or_throw<PrimExpr>(), out);
- return;
- }
+ if (const auto* and_node = pred.as<prim::BitwiseAndNode>()) {
+ FlattenConjuncts(and_node->a, out);
+ FlattenConjuncts(and_node->b, out);
+ return;
}
out->push_back(pred);
}
@@ -1455,17 +1449,13 @@ class TilePrimitiveDispatcher : public StmtExprMutator {
int PushPredicateCtx(const PrimExpr& pred) {
if (ctx_stack_.empty()) return 0;
- if (const auto* and_node = pred.as<prim::AndNode>()) {
- (void)and_node;
+ if (pred.as<prim::AndNode>() || pred.as<prim::BitwiseAndNode>()) {
return PushConjunctivePredicateCtx(pred);
}
if (const auto* call = pred.as<CallNode>()) {
if (call->op.same_as(tirx::builtin::filter())) {
return PushFilterPredicateCtx(call);
}
- if (IsBitwiseAndCall(call)) {
- return PushConjunctivePredicateCtx(pred);
- }
}
if (TryPushComparisonPredicate(pred)) return 1;
return 0;
@@ -1486,6 +1476,41 @@ class TilePrimitiveDispatcher : public StmtExprMutator {
}
return PrimExpr(a && b);
}
+ if (const auto* op = pred.as<prim::LShiftNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ PrimExpr b = RewriteFilterCalls(op->b);
+ if (a.same_as(op->a) && b.same_as(op->b)) return pred;
+ return tvm::left_shift(a, b, op->span);
+ }
+ if (const auto* op = pred.as<prim::RShiftNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ PrimExpr b = RewriteFilterCalls(op->b);
+ if (a.same_as(op->a) && b.same_as(op->b)) return pred;
+ return tvm::right_shift(a, b, op->span);
+ }
+ if (const auto* op = pred.as<prim::BitwiseAndNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ PrimExpr b = RewriteFilterCalls(op->b);
+ if (a.same_as(op->a) && b.same_as(op->b)) return pred;
+ return tvm::bitwise_and(a, b, op->span);
+ }
+ if (const auto* op = pred.as<prim::BitwiseOrNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ PrimExpr b = RewriteFilterCalls(op->b);
+ if (a.same_as(op->a) && b.same_as(op->b)) return pred;
+ return tvm::bitwise_or(a, b, op->span);
+ }
+ if (const auto* op = pred.as<prim::BitwiseXorNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ PrimExpr b = RewriteFilterCalls(op->b);
+ if (a.same_as(op->a) && b.same_as(op->b)) return pred;
+ return tvm::bitwise_xor(a, b, op->span);
+ }
+ if (const auto* op = pred.as<prim::BitwiseNotNode>()) {
+ PrimExpr a = RewriteFilterCalls(op->a);
+ if (a.same_as(op->a)) return pred;
+ return prim::BitwiseNot(a, op->span);
+ }
if (const auto* call = pred.as<CallNode>()) {
if (call->op.same_as(tirx::builtin::filter())) {
return RewriteFilterCalls(RewriteFilterCall(call));
diff --git a/src/tirx/transform/vectorize_loop.cc
b/src/tirx/transform/vectorize_loop.cc
index f79553c581..bed3ac8f44 100644
--- a/src/tirx/transform/vectorize_loop.cc
+++ b/src/tirx/transform/vectorize_loop.cc
@@ -597,6 +597,28 @@ class Vectorizer : public StmtExprMutator {
return BinaryVec<prim::Or>(op, inplace_mode);
}
+ UnchangedOr<PrimExpr> Mutate_(const prim::LShiftNode* op, InplaceMode
inplace_mode) final {
+ return BinaryVec<prim::LShift>(op, inplace_mode);
+ }
+ UnchangedOr<PrimExpr> Mutate_(const prim::RShiftNode* op, InplaceMode
inplace_mode) final {
+ return BinaryVec<prim::RShift>(op, inplace_mode);
+ }
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseAndNode* op, InplaceMode
inplace_mode) final {
+ return BinaryVec<prim::BitwiseAnd>(op, inplace_mode);
+ }
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseOrNode* op, InplaceMode
inplace_mode) final {
+ return BinaryVec<prim::BitwiseOr>(op, inplace_mode);
+ }
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseXorNode* op, InplaceMode
inplace_mode) final {
+ return BinaryVec<prim::BitwiseXor>(op, inplace_mode);
+ }
+
+ UnchangedOr<PrimExpr> Mutate_(const prim::BitwiseNotNode* op, InplaceMode
inplace_mode) final {
+ auto a = this->Mutate(op->a, inplace_mode);
+ if (a.UnchangedOrSameAs(op->a)) return ffi::Unchanged();
+ return prim::BitwiseNot(std::move(a).ValueOrUnchanged(op->a), op->span);
+ }
+
UnchangedOr<PrimExpr> Mutate_(const prim::NotNode* op, InplaceMode
inplace_mode) final {
auto a_update = this->Mutate(op->a, inplace_mode);
bool a_unchanged = a_update.UnchangedOrSameAs(op->a);
@@ -1235,7 +1257,7 @@ class Vectorizer : public StmtExprMutator {
int b_lanes = GetLanesOrVScaleFactor(b.ty());
int lanes = std::max(a_lanes, b_lanes);
bool is_scalable = a.ty().IsScalableVector() ||
b.ty().IsScalableVector();
- return TOp(BroadcastTo(a, lanes, is_scalable), BroadcastTo(b, lanes,
is_scalable));
+ return TOp(BroadcastTo(a, lanes, is_scalable), BroadcastTo(b, lanes,
is_scalable), op->span);
}
}
template <typename T, typename FCompute>
diff --git a/tests/python/sym/test_sym_rewrite_simplify.py
b/tests/python/sym/test_sym_rewrite_simplify.py
index 5236fd0938..61fd4084f0 100644
--- a/tests/python/sym/test_sym_rewrite_simplify.py
+++ b/tests/python/sym/test_sym_rewrite_simplify.py
@@ -1278,7 +1278,7 @@ class TestCast(BaseCompare):
class TestShiftLeft(BaseCompare):
- z = tvm.tirx.op.call_intrin("int32", "prim.shift_left", 1, 10)
+ z = tvm.tirx.LShift(tvm.tirx.const(1, "int32"), tvm.tirx.const(10,
"int32"))
test_case = tvm.testing.parameter(
TestCase(z, tvm.tirx.const(1 << 10, "int32")),
)
diff --git a/tests/python/tirx-base/test_tir_nodes.py
b/tests/python/tirx-base/test_tir_nodes.py
index 08b8ae2bdc..e79f01fb55 100644
--- a/tests/python/tirx-base/test_tir_nodes.py
+++ b/tests/python/tirx-base/test_tir_nodes.py
@@ -199,19 +199,19 @@ def test_all():
def test_bitwise():
x = tvm.tirx.Var("x", "int32")
y = tvm.tirx.Var("y", "int32")
- assert str(x << y) == "T.shift_left(x, y)"
- assert str(x >> y) == "T.shift_right(x, y)"
- assert str(x & y) == "T.bitwise_and(x, y)"
- assert str(x | y) == "T.bitwise_or(x, y)"
- assert str(x ^ y) == "T.bitwise_xor(x, y)"
- assert str(10 & x) == "T.bitwise_and(10, x)"
- assert str(10 | x) == "T.bitwise_or(10, x)"
- assert str(10 ^ x) == "T.bitwise_xor(10, x)"
- assert str(10 >> x) == "T.shift_right(10, x)"
- assert str(10 << x) == "T.shift_left(10, x)"
+ assert str(x << y) == "x << y"
+ assert str(x >> y) == "x >> y"
+ assert str(x & y) == "x & y"
+ assert str(x | y) == "x | y"
+ assert str(x ^ y) == "x ^ y"
+ assert str(10 & x) == "10 & x"
+ assert str(10 | x) == "10 | x"
+ assert str(10 ^ x) == "10 ^ x"
+ assert str(10 >> x) == "10 >> x"
+ assert str(10 << x) == "10 << x"
assert str(10 % x) == "10 % x"
- assert str(~x) == "T.bitwise_not(x)"
+ assert str(~x) == "~x"
assert (tvm.tirx.const(1, "int8x2") >> 1).ty.dtype == "int8x2"
assert (x >> tvm.tirx.const(1, "int32x2")).ty.dtype == "int32x2"
assert (tvm.tirx.Var("z", "int8x2") << tvm.tirx.const(1,
"int8x2")).ty.dtype == "int8x2"
diff --git a/tests/python/tirx-base/test_tir_op_types.py
b/tests/python/tirx-base/test_tir_op_types.py
index eaee54309b..3eb41d56fc 100644
--- a/tests/python/tirx-base/test_tir_op_types.py
+++ b/tests/python/tirx-base/test_tir_op_types.py
@@ -269,27 +269,27 @@ def test_tir_op_shift_left():
x = tirx.Var("x", ty="int32")
y = tirx.Var("x", ty="int32")
expr = tirx.shift_left(x, y)
- assert expr.op.name == "prim.shift_left"
+ assert isinstance(expr, tirx.LShift)
def test_tir_op_shift_right():
x = tirx.Var("x", ty="int32")
y = tirx.Var("x", ty="int32")
expr = tirx.shift_right(x, y)
- assert expr.op.name == "prim.shift_right"
+ assert isinstance(expr, tirx.RShift)
def test_tir_op_bitwise():
x = tirx.Var("x", ty="int32")
y = tirx.Var("y", ty="int32")
expr = tirx.bitwise_and(x, y)
- assert expr.op.name == "prim.bitwise_and"
+ assert isinstance(expr, tirx.BitwiseAnd)
expr = tirx.bitwise_or(x, y)
- assert expr.op.name == "prim.bitwise_or"
+ assert isinstance(expr, tirx.BitwiseOr)
expr = tirx.bitwise_not(x)
- assert expr.op.name == "prim.bitwise_not"
+ assert isinstance(expr, tirx.BitwiseNot)
expr = tirx.bitwise_xor(x, y)
- assert expr.op.name == "prim.bitwise_xor"
+ assert isinstance(expr, tirx.BitwiseXor)
def test_tir_op_TVMBackendAllocWorkspace():
diff --git a/tests/python/tirx-transform/test_tir_transform_vectorize.py
b/tests/python/tirx-transform/test_tir_transform_vectorize.py
index dd029eb7d2..379bbfd4b2 100644
--- a/tests/python/tirx-transform/test_tir_transform_vectorize.py
+++ b/tests/python/tirx-transform/test_tir_transform_vectorize.py
@@ -667,10 +667,7 @@ def test_vectorize_nested_predicates_preserve_both_masks():
tvm_ffi.structural_walk(after.body, (tvm.ir.Call, collect_predicates))
assert len(predicates) == 2
- assert any(
- isinstance(predicate, tvm.ir.Call) and predicate.op.name ==
"prim.bitwise_and"
- for predicate in predicates
- )
+ assert any(isinstance(predicate, tvm.tirx.BitwiseAnd) for predicate in
predicates)
def test_vectorize_and_predicate_invalid_conditions():
diff --git a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py
b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py
index 06e7539951..72de49766a 100644
--- a/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py
+++ b/tests/python/tirx/operator/tile_primitive/cuda/copy/test_reg.py
@@ -835,18 +835,16 @@ def _eval_const_layout_expr(expr, values):
return lhs % rhs
if node_type == "Cast":
return _eval_const_layout_expr(expr.value, values)
- if node_type == "Call":
- args = [_eval_const_layout_expr(arg, values) for arg in expr.args]
- op_name = str(expr.op.name)
- if op_name == "prim.bitwise_xor":
- return args[0] ^ args[1]
- if op_name == "prim.bitwise_and":
- return args[0] & args[1]
- if op_name == "prim.shift_left":
- return args[0] << args[1]
- if op_name == "prim.shift_right":
- return args[0] >> args[1]
- raise AssertionError(f"Cannot evaluate call {op_name}")
+ if node_type in ("BitwiseXor", "BitwiseAnd", "LShift", "RShift"):
+ lhs = _eval_const_layout_expr(expr.a, values)
+ rhs = _eval_const_layout_expr(expr.b, values)
+ if node_type == "BitwiseXor":
+ return lhs ^ rhs
+ if node_type == "BitwiseAnd":
+ return lhs & rhs
+ if node_type == "LShift":
+ return lhs << rhs
+ return lhs >> rhs
raise AssertionError(f"Cannot evaluate node type {node_type}")
diff --git a/tests/python/tirx/test_layout.py b/tests/python/tirx/test_layout.py
index de616be5e4..d7f20f4a1c 100644
--- a/tests/python/tirx/test_layout.py
+++ b/tests/python/tirx/test_layout.py
@@ -1793,18 +1793,16 @@ def _evaluate_layout_expr(expr, values):
return lhs % rhs
if node_type == "Cast":
return _evaluate_layout_expr(expr.value, values)
- if node_type == "Call":
- args = [_evaluate_layout_expr(arg, values) for arg in expr.args]
- op_name = str(expr.op.name)
- if op_name == "prim.bitwise_xor":
- return args[0] ^ args[1]
- if op_name == "prim.bitwise_and":
- return args[0] & args[1]
- if op_name == "prim.shift_left":
- return args[0] << args[1]
- if op_name == "prim.shift_right":
- return args[0] >> args[1]
- raise AssertionError(f"Cannot evaluate call {op_name}")
+ if node_type in ("BitwiseXor", "BitwiseAnd", "LShift", "RShift"):
+ lhs = _evaluate_layout_expr(expr.a, values)
+ rhs = _evaluate_layout_expr(expr.b, values)
+ if node_type == "BitwiseXor":
+ return lhs ^ rhs
+ if node_type == "BitwiseAnd":
+ return lhs & rhs
+ if node_type == "LShift":
+ return lhs << rhs
+ return lhs >> rhs
raise AssertionError(f"Cannot evaluate node type {node_type}")