diff --git a/docs/reference/api/python/script/ir_builder.rst b/docs/reference/api/python/script/ir_builder.rst index 35ff6bfa957f..564519134a1b 100644 --- a/docs/reference/api/python/script/ir_builder.rst +++ b/docs/reference/api/python/script/ir_builder.rst @@ -23,7 +23,7 @@ tvm.script.ir_builder .. automodule:: tvm.script.ir_builder :members: :imported-members: - :exclude-members: Call, DataTypeImm, FuncType, GenericConst, PrimType, Range, StringImm, StringType, Type + :exclude-members: Call, DataTypeImm, FuncType, GenericConst, MissingType, PrimType, Range, StringImm, StringType, Type tvm.relax.script.ir_builder.distributed *************************************** diff --git a/include/tvm/ir/base_expr.h b/include/tvm/ir/base_expr.h index cbda0a4dff4e..9ce81b197884 100644 --- a/include/tvm/ir/base_expr.h +++ b/include/tvm/ir/base_expr.h @@ -79,15 +79,36 @@ class TypeNode : public ffi::Object { */ class Type : public ffi::ObjectRef { public: - /*! \brief Sentinel for a type that has not been populated yet. */ + /*! \brief Construct a MissingType for type information not yet populated. */ TVM_DLL static Type Missing(); - /*! \return whether this is the missing-type sentinel. */ - TVM_DLL bool IsMissing() const; - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Type, ffi::ObjectRef, TypeNode); }; +/*! + * \brief Type information that has not been supplied or computed. + * + * MissingType is not a concrete type, wildcard, or inference variable. Unlike + * AnyType and Void, it must be resolved before a fully typed IR boundary. + */ +class MissingTypeNode final : public TypeNode { + public: + static void RegisterReflection() { + namespace refl = tvm::ffi::reflection; + refl::ObjectDef(); + } + + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("ir.MissingType", MissingTypeNode, TypeNode); +}; + +/*! \brief Managed reference to the MissingTypeNode singleton. */ +class MissingType final : public Type { + public: + TVM_DLL MissingType(); + + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(MissingType, Type, MissingTypeNode); +}; + /*! * \brief Type marker for opaque construction-time expressions. * @@ -311,10 +332,10 @@ class ExprNode : public ffi::Object { /*! * \brief The deduced or annotated type of the expression. * - * Type::Missing() denotes type information that will be populated by + * MissingType() denotes type information that will be populated by * later analysis passes instead of expression constructors. */ - mutable Type ty = Type::Missing(); + mutable Type ty = MissingType(); static void RegisterReflection() { namespace refl = tvm::ffi::reflection; @@ -322,7 +343,7 @@ class ExprNode : public ffi::Object { refl::ObjectDef() .def_ro("span", &ExprNode::span, refl::DefaultValue(Span()), refl::AttachFieldFlag::SEqHashIgnore()) - .def_ro("ty", &ExprNode::ty, refl::DefaultValue(Type::Missing())); + .def_ro("ty", &ExprNode::ty, refl::DefaultValue(MissingType())); } static constexpr TVMFFISEqHashKind _type_s_eq_hash_kind = kTVMFFISEqHashKindTreeNode; diff --git a/include/tvm/relax/expr_functor.h b/include/tvm/relax/expr_functor.h index cc722a6d8dba..e6d305a4e4a1 100644 --- a/include/tvm/relax/expr_functor.h +++ b/include/tvm/relax/expr_functor.h @@ -524,7 +524,7 @@ class ExprMutatorBase : public ExprFunctor { bool VisitAndCheckTypeFieldUnchanged(const ffi::ObjectRef& ty) { if (const TypeNode* ty_node = ty.as()) { Type type = ffi::GetRef(ty_node); - return type.IsMissing() || this->VisitExprDepTypeField(type).same_as(ty); + return type.as().has_value() || this->VisitExprDepTypeField(type).same_as(ty); } else { return true; } diff --git a/include/tvm/relax/type.h b/include/tvm/relax/type.h index 5b3b65202473..8e4fb36e9e69 100644 --- a/include/tvm/relax/type.h +++ b/include/tvm/relax/type.h @@ -175,7 +175,7 @@ class TensorTypeNode : public TypeNode { ffi::Optional> GetShape() const { if (!shape.has_value()) return {}; const Expr& shape_expr = this->shape.value(); - if (shape_expr->ty.IsMissing()) return {}; + if (shape_expr->ty.as().has_value()) return {}; if (const auto* shape_ty = shape_expr->ty.as()) { return shape_ty->values; } @@ -366,7 +366,7 @@ inline ffi::Optional MatchType(const Expr& expr) { */ template inline const T* GetTypeAs(const Expr& expr) { - TVM_FFI_ICHECK(!expr->ty.IsMissing()) + TVM_FFI_ICHECK(!expr->ty.as().has_value()) << "The type is not populated, check if you have normalized the expr"; return expr->ty.as(); } @@ -378,7 +378,7 @@ inline const T* GetTypeAs(const Expr& expr) { * \return underlying Relax type. */ inline Type GetType(const Expr& expr) { - TVM_FFI_ICHECK(!expr->ty.IsMissing()) + TVM_FFI_ICHECK(!expr->ty.as().has_value()) << "The type is not populated, check if you have normalized the expr"; return expr->ty; } diff --git a/include/tvm/relax/type_functor.h b/include/tvm/relax/type_functor.h index b1997e0709b5..5ce7f7c61a17 100644 --- a/include/tvm/relax/type_functor.h +++ b/include/tvm/relax/type_functor.h @@ -73,8 +73,7 @@ class TypeFunctor { */ virtual R VisitType(const Type& n, Args... args) { TVM_FFI_ICHECK(n.defined()); - TVM_FFI_ICHECK_NE(n->type_index(), TypeNode::RuntimeTypeIndex()) - << "TypeFunctor cannot visit Type::Missing()"; + TVM_FFI_ICHECK(!n.as().has_value()) << "TypeFunctor cannot visit Type::Missing()"; static FType vtable = InitVTable(); return vtable(n, this, std::forward(args)...); } diff --git a/include/tvm/tirx/op.h b/include/tvm/tirx/op.h index 32b417ba2ee4..094083381025 100644 --- a/include/tvm/tirx/op.h +++ b/include/tvm/tirx/op.h @@ -338,7 +338,7 @@ TVM_DECLARE_INTRIN_BINARY(ldexp); * \return The check results */ inline bool IsPointerType(const Type& type, DLDataType element_type) { - if (type.IsMissing()) return false; + if (type.as().has_value()) return false; if (const auto* ptr_type = type.as()) { if (const auto* prim_type = ptr_type->element_type.as()) { return prim_type->dtype == element_type; diff --git a/python/tvm/ir/__init__.py b/python/tvm/ir/__init__.py index fe264b16ba8a..7703a6dc4043 100644 --- a/python/tvm/ir/__init__.py +++ b/python/tvm/ir/__init__.py @@ -33,7 +33,17 @@ # Register Type before Expr. Expr's reflected ``ty`` field otherwise creates # an auto-generated Type wrapper before the concrete Python class is available. -from .type import AnyType, FuncType, OpaqueType, PointerType, PrimType, StringType, TupleType, Type +from .type import ( + AnyType, + FuncType, + MissingType, + OpaqueType, + PointerType, + PrimType, + StringType, + TupleType, + Type, +) from .expr import ( Call, Constant, diff --git a/python/tvm/ir/expr.py b/python/tvm/ir/expr.py index 177029702ece..1177f52c5631 100644 --- a/python/tvm/ir/expr.py +++ b/python/tvm/ir/expr.py @@ -48,7 +48,7 @@ class Expr(Node): ty: "tvm.ir.Type" def __getitem__(self, index): - if self.ty.is_missing(): + if isinstance(self.ty, tvm.ir.MissingType): # Preserve Relax's pre-normalization tuple access: operator calls # have a missing result type until the block builder infers it. return TupleGetItem(self, index) diff --git a/python/tvm/ir/type.py b/python/tvm/ir/type.py index e0fbf6cd099a..54172b637fb3 100644 --- a/python/tvm/ir/type.py +++ b/python/tvm/ir/type.py @@ -30,18 +30,14 @@ class Type(Node, Scriptable): @staticmethod def missing(): - """Return the sentinel for missing type information.""" + """Construct a MissingType for missing type information.""" return _ffi_api.TypeMissing() @staticmethod def Missing(): - """Return the sentinel for missing type information.""" + """Construct a MissingType for missing type information.""" return _ffi_api.TypeMissing() - def is_missing(self): - """Return whether this is the missing-type sentinel.""" - return _ffi_api.TypeIsMissing(self) - def __eq__(self, other): """Compare two types for structural equivalence.""" return bool(tvm_ffi.structural_equal(self, other)) @@ -54,6 +50,18 @@ def same_as(self, other): return self.is_(other) +@tvm_ffi.register_object("ir.MissingType") +class MissingType(Type): + """Type information that has not been supplied or computed. + + Unlike AnyType or Void, this is not a concrete type and must be resolved + before a boundary that requires fully typed IR. + """ + + def __init__(self): + self.__init_handle_by_constructor__(_ffi_api.MissingType) + + @tvm_ffi.register_object("ir.AnyType") class AnyType(Type): """The top type, which admits any value.""" diff --git a/python/tvm/relax/expr.py b/python/tvm/relax/expr.py index 0c3a1f579ffc..917dcfd03133 100644 --- a/python/tvm/relax/expr.py +++ b/python/tvm/relax/expr.py @@ -89,7 +89,7 @@ def _relax_type_is_base_of(self: Type, derived: Type) -> bool: def _is_tensor_or_missing_type(ty: Type) -> bool: - return isinstance(ty, tvm.relax.TensorType) or ty.is_missing() + return isinstance(ty, tvm.relax.TensorType | tvm.ir.MissingType) def _binary_op_helper(lhs: Expr, rhs: Expr, op: Callable): diff --git a/python/tvm/relax/frontend/nn/subroutine.py b/python/tvm/relax/frontend/nn/subroutine.py index a589919b9f41..b3bb6138ab49 100644 --- a/python/tvm/relax/frontend/nn/subroutine.py +++ b/python/tvm/relax/frontend/nn/subroutine.py @@ -45,7 +45,7 @@ def _normalize_expr(block_builder, arg, as_relax_expr=False): if isinstance(arg, tuple): arg = relax.Tuple([_normalize_expr(block_builder, element) for element in arg]) - if isinstance(arg, relax.Expr) and arg.ty.is_missing(): + if isinstance(arg, relax.Expr) and isinstance(arg.ty, ir.MissingType): arg = block_builder.emit(arg) if isinstance(arg, nn.Tensor) and as_relax_expr: @@ -108,7 +108,7 @@ def new_forward(self, *args, **kwargs): out = subroutine(*subroutine_args) if is_nn_tensor_output: - if out.ty.is_missing(): + if isinstance(out.ty, ir.MissingType): out = block_builder.emit(out, name_hint=f"{subroutine.name_hint}_output") out = nn.Tensor(_expr=out) return out diff --git a/python/tvm/relax/op/__init__.py b/python/tvm/relax/op/__init__.py index d8628c806d0c..f1035fe43e19 100644 --- a/python/tvm/relax/op/__init__.py +++ b/python/tvm/relax/op/__init__.py @@ -193,9 +193,8 @@ def _unary(lhs, op): return op(lhs) def _call(func, *args, attrs=None): - if not ( - isinstance(func.ty, expr.tvm.ir.FuncType | expr.tvm.relax.FuncType) - or func.ty.is_missing() + if not isinstance( + func.ty, expr.tvm.ir.FuncType | expr.tvm.relax.FuncType | expr.tvm.ir.MissingType ): return NotImplemented return expr.tvm.ir.Call(func, args, attrs=attrs) diff --git a/python/tvm/relax/testing/ast_printer.py b/python/tvm/relax/testing/ast_printer.py index 9a467eb577d5..0ffdb74d96ec 100644 --- a/python/tvm/relax/testing/ast_printer.py +++ b/python/tvm/relax/testing/ast_printer.py @@ -91,7 +91,7 @@ def build_expr(self, node: relax.Expr, nodename: str, force_newline=False, **kwa Handles whether to include the ty fields. """ fields = kwargs.copy() - if not node.ty.is_missing() and self.include_ty_annotations: + if not isinstance(node.ty, tvm.ir.MissingType) and self.include_ty_annotations: fields["ty"] = self.visit_ty_(node.ty) return self.build_ast_node(nodename, force_newline=force_newline, **fields) diff --git a/python/tvm/script/ir_builder/__init__.py b/python/tvm/script/ir_builder/__init__.py index 9e1f1f7f0fd9..30f92ab098e1 100644 --- a/python/tvm/script/ir_builder/__init__.py +++ b/python/tvm/script/ir_builder/__init__.py @@ -21,6 +21,7 @@ DataTypeImm, FuncType, GenericConst, + MissingType, PrimType, Range, StringImm, @@ -67,6 +68,7 @@ "GenericConst", "IRBuilder", "IRModuleFrame", + "MissingType", "PrimType", "Range", "StringImm", diff --git a/src/ir/expr.cc b/src/ir/expr.cc index 3399805126b1..7d07ab95d90a 100644 --- a/src/ir/expr.cc +++ b/src/ir/expr.cc @@ -725,7 +725,7 @@ Tuple::Tuple(ffi::Array fields, Span span) : Expr(ffi::UnsafeInit{}) { ffi::Optional tuple_ty = [&]() -> ffi::Optional { ffi::Array field_ty; for (const Expr& field : fields) { - if (field->ty.IsMissing()) { + if (field->ty.as().has_value()) { return std::nullopt; } field_ty.push_back(field->ty); @@ -821,7 +821,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { // Constants GenericConst::GenericConst(ffi::Any value, Type ty, Span span) : Constant(ffi::UnsafeInit{}) { - TVM_FFI_CHECK(!ty.IsMissing(), TypeError) << "GenericConst requires an expression type"; + TVM_FFI_CHECK(!ty.as().has_value(), TypeError) + << "GenericConst requires an expression type"; TVM_FFI_CHECK(!value.as() && !value.as() && !value.as() && !value.as(), TypeError) @@ -1217,7 +1218,7 @@ Type Call::ReinferType(const CallNode* call) { TVM_FFI_CHECK(infer_type.count(op.value()), ValueError) << "No context-free FInferType hook is registered for " << op.value(); Type result = infer_type[op.value()].CallExpected(call).value(); - TVM_FFI_CHECK(!result.IsMissing(), InternalError) + TVM_FFI_CHECK(!result.as().has_value(), InternalError) << "FInferType for " << op.value() << " returned Type::Missing()"; return result; } diff --git a/src/ir/prim/const_fold.h b/src/ir/prim/const_fold.h index 77d42f09e990..0a0f611ec7e2 100644 --- a/src/ir/prim/const_fold.h +++ b/src/ir/prim/const_fold.h @@ -79,7 +79,7 @@ inline bool IsIndexType(DLDataType type) { inline bool IsIndexTypedExpr(const ExprNode* expr) { TVM_FFI_DCHECK(expr != nullptr); - TVM_FFI_DCHECK(!expr->ExprNode::ty.IsMissing()); + TVM_FFI_DCHECK(!expr->ExprNode::ty.as().has_value()); const auto* prim_ty = expr->ExprNode::ty.as(); TVM_FFI_DCHECK(prim_ty != nullptr); return IsIndexType(prim_ty->dtype); diff --git a/src/ir/prim/expr.cc b/src/ir/prim/expr.cc index 8a64472d7adc..295d0b77455a 100644 --- a/src/ir/prim/expr.cc +++ b/src/ir/prim/expr.cc @@ -40,7 +40,7 @@ int GetLanesOrVScaleFactor(const PrimType& ty) { TVM_FFI_INLINE const PrimTypeNode* GetPrimTypeNode(const PrimExpr& expr) { const auto* node = expr.get(); TVM_FFI_DCHECK(node != nullptr); - TVM_FFI_DCHECK(!node->ExprNode::ty.IsMissing()); + TVM_FFI_DCHECK(!node->ExprNode::ty.as().has_value()); const auto* prim_ty = node->ExprNode::ty.as(); TVM_FFI_DCHECK(prim_ty != nullptr); return prim_ty; diff --git a/src/ir/prim/op.cc b/src/ir/prim/op.cc index 39be7bb5a895..bdbdba9d0664 100644 --- a/src/ir/prim/op.cc +++ b/src/ir/prim/op.cc @@ -34,7 +34,7 @@ TVM_FFI_INLINE const PrimTypeNode* GetPrimTypeNode(const PrimExpr& expr) { // Avoid PrimExpr::ty() ObjectRef materialization on binary operator hot paths. const auto* node = expr.get(); TVM_FFI_DCHECK(node != nullptr); - TVM_FFI_DCHECK(!node->ExprNode::ty.IsMissing()); + TVM_FFI_DCHECK(!node->ExprNode::ty.as().has_value()); const auto* prim_ty = node->ExprNode::ty.as(); TVM_FFI_DCHECK(prim_ty != nullptr); return prim_ty; diff --git a/src/ir/type.cc b/src/ir/type.cc index e824b4a3807f..dfa2601f07e9 100644 --- a/src/ir/type.cc +++ b/src/ir/type.cc @@ -63,7 +63,7 @@ ffi::ObjectPtr GetCachedPrimTypeNode(DLDataType dtype) { TVM_FFI_INLINE ffi::Expected> TypeVisit( ffi::StructuralVisitorObj*, ffi::AnyView) noexcept { - // Type::Missing() is the only concrete TypeNode value; span is ignored debug metadata. + // Field-less types are leaves; span is ignored debug metadata. return std::nullopt; } @@ -246,14 +246,7 @@ TVM_FFI_INLINE ffi::Expected> TupleTypeMaybeInplaceMu } // namespace -Type Type::Missing() { - static Type missing = []() { - Type type(ffi::UnsafeInit{}); - type.data_ = ffi::make_object(); - return type; - }(); - return missing; -} +Type Type::Missing() { return MissingType(); } TVM_FFI_STATIC_INIT_BLOCK() { namespace refl = tvm::ffi::reflection; @@ -263,12 +256,24 @@ TVM_FFI_STATIC_INIT_BLOCK() { .attr(refl::type_attr::kStructuralMutate, ffi::FStructuralMutate::FromNative<&TypeMutate>()) .attr(refl::type_attr::kStructuralMaybeInplaceMutate, ffi::FStructuralMutate::FromNative<&TypeMaybeInplaceMutate>()); - refl::GlobalDef() - .def("ir.TypeMissing", []() { return Type::Missing(); }) - .def("ir.TypeIsMissing", [](Type type) { return type.IsMissing(); }); + refl::GlobalDef().def("ir.TypeMissing", []() { return Type::Missing(); }); +} + +MissingType::MissingType() : Type(ffi::UnsafeInit{}) { + static const auto singleton = ffi::make_object(); + data_ = singleton; } -bool Type::IsMissing() const { return this->same_as(Type::Missing()); } +TVM_FFI_STATIC_INIT_BLOCK() { + namespace refl = tvm::ffi::reflection; + MissingTypeNode::RegisterReflection(); + refl::TypeAttrDef() + .attr(refl::type_attr::kStructuralVisit, ffi::FStructuralVisit::FromNative<&TypeVisit>()) + .attr(refl::type_attr::kStructuralMutate, ffi::FStructuralMutate::FromNative<&TypeMutate>()) + .attr(refl::type_attr::kStructuralMaybeInplaceMutate, + ffi::FStructuralMutate::FromNative<&TypeMaybeInplaceMutate>()); + refl::GlobalDef().def("ir.MissingType", []() { return MissingType(); }); +} AnyType::AnyType(Span span) : Type(ffi::UnsafeInit{}) { ffi::ObjectPtr n = ffi::make_object(); @@ -390,7 +395,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { // PointerType PointerType::PointerType(Type element_type, ffi::String storage_scope) : Type(ffi::UnsafeInit{}) { - TVM_FFI_ICHECK(!element_type.IsMissing()) << "PointerType element_type cannot be Type::Missing()"; + TVM_FFI_ICHECK(!element_type.as().has_value()) + << "PointerType element_type cannot be Type::Missing()"; ffi::ObjectPtr n = ffi::make_object(); if (storage_scope.empty()) { n->storage_scope = "global"; diff --git a/src/relax/analysis/type_analysis.cc b/src/relax/analysis/type_analysis.cc index a60ffbab82bc..42ffa76f0b63 100644 --- a/src/relax/analysis/type_analysis.cc +++ b/src/relax/analysis/type_analysis.cc @@ -307,6 +307,8 @@ class TypeBaseChecker : public TypeFunctor().has_value() && !other.as().has_value()) + << "Type analysis requires populated types"; // quick path // Note: subclass may disable this quick path if we need to go over all type. if (lhs.same_as(other)) return BaseCheckResult::kPass; @@ -632,6 +634,8 @@ class TypeBasePreconditionCollector : public TypeFunctor().has_value() && !other.as().has_value()) + << "Type analysis requires populated types"; if (lhs.same_as(other)) { // Early bail-out if the Type has reference equality. return IntImm::Bool(true); @@ -990,6 +994,8 @@ class TypeLCAFinder : public TypeFunctor { explicit TypeLCAFinder(sym::AnalyzerObj* ana) : analyzer_(ana) {} Type VisitType(const Type& lhs, const Type& other) final { + TVM_FFI_ICHECK(!lhs.as().has_value() && !other.as().has_value()) + << "Type analysis requires populated types"; // quick path if (lhs.same_as(other)) return lhs; return TypeFunctor::VisitType(lhs, other); diff --git a/src/relax/analysis/well_formed.cc b/src/relax/analysis/well_formed.cc index 6e1965b85a37..b5113b9365b5 100644 --- a/src/relax/analysis/well_formed.cc +++ b/src/relax/analysis/well_formed.cc @@ -186,7 +186,7 @@ class WellFormedChecker : public relax::ExprVisitor, public relax::TypeVisitor { } void VisitExpr(const Expr& expr) final { - if (!expr.as() && expr->ty.IsMissing()) { + if (!expr.as() && expr->ty.as().has_value()) { TVM_FFI_VISIT_THROW(TypeError, expr) << "The ty of Expr " << expr << " is missing."; } relax::ExprVisitor::VisitExpr(expr); @@ -202,7 +202,7 @@ class WellFormedChecker : public relax::ExprVisitor, public relax::TypeVisitor { } } - if (!op->ty.IsMissing()) { + if (!op->ty.as().has_value()) { if (!op->ty->IsInstance()) { TVM_FFI_VISIT_THROW(TypeError, var) << "The ty of GlobalVar " << ffi::GetRef(op) << " must be either FuncType."; @@ -323,7 +323,7 @@ class WellFormedChecker : public relax::ExprVisitor, public relax::TypeVisitor { CheckType(param.get()); } // check function ret_ty - if (!op->ret_ty.IsMissing()) { + if (!op->ret_ty.as().has_value()) { this->VisitType(op->ret_ty); } else { TVM_FFI_VISIT_THROW(TypeError, ffi::GetRef(op)) << "Function must have defined ret_ty"; @@ -428,7 +428,8 @@ class WellFormedChecker : public relax::ExprVisitor, public relax::TypeVisitor { has_infer_type = op_map_infer_type_.count(op.value()) || op_map_infer_type_with_builder_.count(op.value()); } - if (check_ty && !call->ty.IsMissing() && (!call->ty.as() || has_infer_type)) { + if (check_ty && !call->ty.as().has_value() && + (!call->ty.as() || has_infer_type)) { // The `InferType` method isn't currently exposed by the // Normalizer, and can only be called indirectly by normalizing // an expression that does not yet have `Type`. @@ -544,7 +545,8 @@ class WellFormedChecker : public relax::ExprVisitor, public relax::TypeVisitor { this->VisitVarDef(binding->var); - if (check_ty && !binding->var->ty.IsMissing() && !binding->value->ty.IsMissing()) { + if (check_ty && !binding->var->ty.as().has_value() && + !binding->value->ty.as().has_value()) { auto expr_ty = GetType(binding->value); auto var_ty = GetType(binding->var); if (!IsBaseOf(var_ty, expr_ty)) { diff --git a/src/relax/backend/vm/lower_runtime_builtin.cc b/src/relax/backend/vm/lower_runtime_builtin.cc index 9872cb81f8dd..93012d950e46 100644 --- a/src/relax/backend/vm/lower_runtime_builtin.cc +++ b/src/relax/backend/vm/lower_runtime_builtin.cc @@ -125,7 +125,7 @@ class LowerRuntimeBuiltinMutator : public ExprMutator { Expr Reshape(const Call& call_node) { TVM_FFI_ICHECK(call_node->args.size() == 2); - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); auto arg = call_node->args[1]; TVM_FFI_CHECK(arg->ty->IsInstance(), TypeError) @@ -140,14 +140,14 @@ class LowerRuntimeBuiltinMutator : public ExprMutator { Expr ShapeOf(const Call& call_node) { TVM_FFI_ICHECK(call_node->args.size() == 1); - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); return Call::Unchecked(Type::Missing(), builtin_shape_of_, call_node->args, Attrs(), {GetType(call_node)}); } Expr TensorToShape(const Call& call_node) { TVM_FFI_ICHECK(call_node->args.size() == 1); - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); return Call::Unchecked(Type::Missing(), builtin_tensor_to_shape_, call_node->args, Attrs(), {GetType(call_node)}); @@ -155,7 +155,7 @@ class LowerRuntimeBuiltinMutator : public ExprMutator { Expr CallPyFunc(const Call& call_node) { TVM_FFI_ICHECK(call_node->args.size() == 2); - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); // Create tuple with function name and arguments tuple ffi::Array tuple_fields; @@ -171,7 +171,7 @@ class LowerRuntimeBuiltinMutator : public ExprMutator { Expr ToDevice(const Call& call_node) { // TODO(yongwww): replace ToVDeviceAttrs with related Expr TVM_FFI_ICHECK(call_node->args.size() == 1); - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); auto attrs = call_node->attrs.as(); ffi::Array args; args.push_back(call_node->args[0]); diff --git a/src/relax/distributed/transform/lower_distir.cc b/src/relax/distributed/transform/lower_distir.cc index 121f546a94a2..4bc9d35b3820 100644 --- a/src/relax/distributed/transform/lower_distir.cc +++ b/src/relax/distributed/transform/lower_distir.cc @@ -114,7 +114,7 @@ class DistIRSharder : public ExprMutator { } Expr ShardInputParamTensorAndConstant(Expr input) { - TVM_FFI_ICHECK(!input->ty.IsMissing()); + TVM_FFI_ICHECK(!input->ty.as().has_value()); Type old_ty = GetType(input); Type new_ty = ConvertType(old_ty, false); if (const auto* var = input.as()) { diff --git a/src/relax/distributed/transform/utils.cc b/src/relax/distributed/transform/utils.cc index 586ac9efeca1..df93811ca14b 100644 --- a/src/relax/distributed/transform/utils.cc +++ b/src/relax/distributed/transform/utils.cc @@ -48,7 +48,7 @@ bool TypeCompatibleWithRelax(ffi::Array tys) { bool IsDistIRFunc(Function func) { ffi::Array param_tys; for (const auto& param : func->params) { - TVM_FFI_ICHECK(!param->ty.IsMissing()); + TVM_FFI_ICHECK(!param->ty.as().has_value()); param_tys.push_back(param->ty.as_or_throw()); } bool compatible_with_dist_ir = TypeCompatibleWithDistIR(param_tys); diff --git a/src/relax/ir/block_builder.cc b/src/relax/ir/block_builder.cc index cf0293b5f41c..af635a71fc4f 100644 --- a/src/relax/ir/block_builder.cc +++ b/src/relax/ir/block_builder.cc @@ -90,7 +90,7 @@ class BlockBuilderImpl : public BlockBuilderNode { GlobalVar gvar(func_name); Type finfo = Type::Missing(); - if (!func->ty.IsMissing()) { + if (!func->ty.as().has_value()) { finfo = GetType(func); } else if (auto* prim_func = func.as()) { // NOTE: use a slightly different type than checked type @@ -275,8 +275,8 @@ class BlockBuilderImpl : public BlockBuilderNode { << "Cannot emit dataflow var in non-dataflow block"; } // normalized check - TVM_FFI_ICHECK(!var_binding->var->ty.IsMissing()); - TVM_FFI_ICHECK(!var_binding->value->ty.IsMissing()); + TVM_FFI_ICHECK(!var_binding->var->ty.as().has_value()); + TVM_FFI_ICHECK(!var_binding->value->ty.as().has_value()); cur_frame->bindings.push_back(binding); binding_table_.insert_or_assign(var_binding->var, var_binding->value); } else if (const auto* match_cast = binding.as()) { @@ -285,8 +285,8 @@ class BlockBuilderImpl : public BlockBuilderNode { << "Cannot emit dataflow var in non-dataflow block"; } // normalized check - TVM_FFI_ICHECK(!match_cast->var->ty.IsMissing()); - TVM_FFI_ICHECK(!match_cast->value->ty.IsMissing()); + TVM_FFI_ICHECK(!match_cast->var->ty.as().has_value()); + TVM_FFI_ICHECK(!match_cast->value->ty.as().has_value()); // NOTE match shape do not follow simple binding rule // as a result should not appear in binding table. cur_frame->bindings.push_back(binding); @@ -532,7 +532,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorIsInstance()) { - TVM_FFI_ICHECK(!normalized->ty.IsMissing()) + TVM_FFI_ICHECK(!normalized->ty.as().has_value()) << "The ty of an Expr except OpNode after " "normalization must not be missing. However, this Expr does not have ty: " << normalized; @@ -590,7 +590,8 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorty.IsMissing()) << "Var " << var->name << " does not have type."; + TVM_FFI_ICHECK(!var->ty.template as().has_value()) + << "Var " << var->name << " does not have type."; return ffi::GetRef(var); } @@ -629,7 +630,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctor(op) : Tuple(new_fields, op->span); // Update tuple fields. - if (tuple->ty.IsMissing()) { + if (tuple->ty.as().has_value()) { ffi::Array tuple_ty; for (Expr field : tuple->fields) { tuple_ty.push_back(GetType(field)); @@ -663,7 +664,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorty.IsMissing()) { + if (call->ty.as().has_value()) { auto inferred_ty = InferType(call); UpdateType(call, inferred_ty); } @@ -732,7 +733,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorty.IsMissing()) { + if (seq_expr->ty.as().has_value()) { UpdateType(seq_expr, EraseToWellDefinedInScope(GetType(seq_expr->body))); } return seq_expr; @@ -751,7 +752,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorspan); } }(); - if (if_node->ty.IsMissing()) { + if (if_node->ty.as().has_value()) { auto true_info = EraseToWellDefinedInScope(GetType(new_true)); auto false_info = EraseToWellDefinedInScope(GetType(new_false)); UpdateType(if_node, TypeLCA(true_info, false_info)); @@ -765,7 +766,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctortuple) ? ffi::GetRef(op) : TupleGetItem(new_tuple, op->index); - if (node->ty.IsMissing()) { + if (node->ty.as().has_value()) { auto opt = MatchType(node->tuple); TVM_FFI_ICHECK(opt) << "The type of Tuple must be TupleType, " << "but expression " << node->tuple << " has type " << node->tuple->ty; @@ -790,7 +791,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorvalue)) { binding = VarBinding(binding->var, new_value, binding->span); } - if (binding->var->ty.IsMissing()) { + if (binding->var->ty.as().has_value()) { UpdateType(binding->var, GetType(new_value)); } return binding; @@ -801,7 +802,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctorvalue)) { binding = MatchCast(binding->var, new_value, binding->ty, binding->span); } - if (binding->var->ty.IsMissing()) { + if (binding->var->ty.as().has_value()) { UpdateType(binding->var, binding->ty); } return binding; @@ -865,7 +866,7 @@ class Normalizer : public BlockBuilderImpl, private ExprFunctor(this)); } else { // derive using function parameters - TVM_FFI_ICHECK(!call->op->ty.IsMissing()); + TVM_FFI_ICHECK(!call->op->ty.as().has_value()); auto opt = MatchType(call->op); TVM_FFI_ICHECK(opt) << "Call->op must contains a function type"; FuncType finfo = opt.value(); diff --git a/src/relax/ir/dependent_type.cc b/src/relax/ir/dependent_type.cc index 6cc4070b35d9..231b3c9f4fac 100644 --- a/src/relax/ir/dependent_type.cc +++ b/src/relax/ir/dependent_type.cc @@ -365,11 +365,12 @@ TVM_FFI_STATIC_INIT_BLOCK() { // Helper functions void UpdateType(Expr expr, Type ty) { - TVM_FFI_ICHECK(expr->ty.IsMissing()) << "To ensure idempotency, " - << "the expression passed to UpdateType " - << "must not have any prior type. " - << "However, expression " << expr << " has type " << expr->ty - << ", which cannot be overwritten with " << ty; + TVM_FFI_ICHECK(expr->ty.as().has_value()) + << "To ensure idempotency, " + << "the expression passed to UpdateType " + << "must not have any prior type. " + << "However, expression " << expr << " has type " << expr->ty + << ", which cannot be overwritten with " << ty; expr->ty = ty; } diff --git a/src/relax/ir/emit_te.cc b/src/relax/ir/emit_te.cc index 24c3d2c86ad3..5880a41cd986 100644 --- a/src/relax/ir/emit_te.cc +++ b/src/relax/ir/emit_te.cc @@ -54,7 +54,8 @@ te::Tensor TETensor(Expr value, ffi::Map tir_var_map, std:: n->shape = std::move(shape); return te::PlaceholderOp(n).output(0); } - TVM_FFI_ICHECK(!value->ty.IsMissing()) << "value must be normalized and contain Type"; + TVM_FFI_ICHECK(!value->ty.as().has_value()) + << "value must be normalized and contain Type"; auto* tensor_ty = GetTypeAs(value); TVM_FFI_ICHECK(tensor_ty) << "Value must be a tensor"; auto* shape_expr = tensor_ty->shape.as(); diff --git a/src/relax/ir/expr.cc b/src/relax/ir/expr.cc index b07914edce4d..1b28fc4e5e7f 100644 --- a/src/relax/ir/expr.cc +++ b/src/relax/ir/expr.cc @@ -647,13 +647,14 @@ Function::Function(ffi::Array params, Expr body, ffi::Optional ret_ty ffi::Array param_ty; for (const Var& param : params) { - TVM_FFI_ICHECK(!param->ty.IsMissing()) << "relax.Function requires params to contain ty"; + TVM_FFI_ICHECK(!param->ty.as().has_value()) + << "relax.Function requires params to contain ty"; param_ty.push_back(GetType(param)); } ffi::Optional body_ty; - if (!body->ty.IsMissing()) { + if (!body->ty.as().has_value()) { body_ty = GetType(body); } @@ -720,7 +721,8 @@ Function Function::CreateEmpty(ffi::Array params, Type ret_ty, bool is_pure Span span) { ffi::Array param_ty; for (const Var& param : params) { - TVM_FFI_ICHECK(!param->ty.IsMissing()) << "relax.Function requires params to contain ty."; + TVM_FFI_ICHECK(!param->ty.as().has_value()) + << "relax.Function requires params to contain ty."; param_ty.push_back(GetType(param)); } @@ -817,7 +819,8 @@ TVM_FFI_STATIC_INIT_BLOCK() { Expr GetShapeOf(const Expr& expr) { // default case, to be normalized. - TVM_FFI_ICHECK(!expr->ty.IsMissing()) << "GetShapeOf can only be applied to normalized expr"; + TVM_FFI_ICHECK(!expr->ty.as().has_value()) + << "GetShapeOf can only be applied to normalized expr"; auto* tinfo = GetTypeAs(expr); TVM_FFI_ICHECK(tinfo != nullptr) << "ShapeOf can only be applied to expr with TensorType"; diff --git a/src/relax/ir/expr_functor.cc b/src/relax/ir/expr_functor.cc index f9e367946708..f4679abf3492 100644 --- a/src/relax/ir/expr_functor.cc +++ b/src/relax/ir/expr_functor.cc @@ -113,7 +113,7 @@ void ExprVisitor::DefaultTypeFieldVisitor::VisitType_(const FuncTypeNode* op) { } void VisitExprDepTypeFieldIfNeeded(ExprVisitor* visitor, const Type& ty) { - if (!ty.IsMissing()) { + if (!ty.as().has_value()) { auto* ty_node = ty.as(); TVM_FFI_DCHECK(ty_node != nullptr); visitor->VisitExprDepTypeField(ffi::GetRef(ty_node)); @@ -521,7 +521,7 @@ Expr ExprMutatorBase::VisitExpr_(const CallNode* call_node) { } Type ret_ty = call_node->ty; - if (!ret_ty.IsMissing()) { + if (!ret_ty.as().has_value()) { ret_ty = this->VisitExprDepTypeField(ret_ty); } bool ret_ty_unchanged = ret_ty.same_as(call_node->ty); @@ -1048,10 +1048,10 @@ ffi::Optional ExprMutator::LookupBinding(const Var& var) { } Var ExprMutator::WithType(Var var, Type ty) { - TVM_FFI_ICHECK(!ty.IsMissing()); + TVM_FFI_ICHECK(!ty.as().has_value()); // TODO(relax-team) add TypeEqual check - if (!var->ty.IsMissing()) { + if (!var->ty.as().has_value()) { // use same-as as a quick path if (var->ty.same_as(ty) || ffi::StructuralEqual()(var->ty, ty)) { return var; diff --git a/src/relax/op/op.cc b/src/relax/op/op.cc index 29b268250dc2..52574c849212 100644 --- a/src/relax/op/op.cc +++ b/src/relax/op/op.cc @@ -409,7 +409,7 @@ ffi::Optional InferCallTIROutputTypeFromArguments( dummy_callee_ty, Call::Unchecked(Type::Missing(), Var("dummy_callee", dummy_callee_ty), dummy_args), BlockBuilder::Create(std::nullopt)); - if (derived_ret_ty.IsMissing()) { + if (derived_ret_ty.as().has_value()) { return std::nullopt; } @@ -1106,7 +1106,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { Type ReturnTensorToShapeType(const CallNode* call_node) { const Call call = ffi::GetRef(call_node); TVM_FFI_ICHECK(call->args.size() == 1); - TVM_FFI_ICHECK(!call->args[0]->ty.IsMissing()); + TVM_FFI_ICHECK(!call->args[0]->ty.as().has_value()); const auto* tensor_ty = GetTypeAs(call->args[0]); TVM_FFI_ICHECK(tensor_ty); TVM_FFI_ICHECK_EQ(tensor_ty->ndim, 1) @@ -1144,7 +1144,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { Type ReturnShapeToTensorType(const CallNode* call_node) { const Call call = ffi::GetRef(call_node); TVM_FFI_ICHECK(call->args.size() == 1); - TVM_FFI_ICHECK(!call->args[0]->ty.IsMissing()); + TVM_FFI_ICHECK(!call->args[0]->ty.as().has_value()); const auto* ty = GetTypeAs(call->args[0]); TVM_FFI_ICHECK(ty); int32_t ndim = ty->ndim; @@ -1505,7 +1505,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { Type InferToVDeviceType(const CallNode* call_node) { const Call call = ffi::GetRef(call_node); TVM_FFI_ICHECK(call->args.size() == 1); - TVM_FFI_ICHECK(!call->args[0]->ty.IsMissing()); + TVM_FFI_ICHECK(!call->args[0]->ty.as().has_value()); TensorType data_ty = GetUnaryInputTensorType(call); auto attrs = call->attrs.as(); VDevice vdev = attrs->dst_vdevice; @@ -1540,7 +1540,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { Type InferHintOnDeviceType(const CallNode* call_node) { const Call call = ffi::GetRef(call_node); TVM_FFI_ICHECK(call->args.size() == 1); - TVM_FFI_ICHECK(!call->args[0]->ty.IsMissing()); + TVM_FFI_ICHECK(!call->args[0]->ty.as().has_value()); TensorType data_ty = GetUnaryInputTensorType(call); return data_ty; } diff --git a/src/relax/op/op_common.h b/src/relax/op/op_common.h index 8d25435fe8f6..2b7c997aa6ad 100644 --- a/src/relax/op/op_common.h +++ b/src/relax/op/op_common.h @@ -115,7 +115,7 @@ namespace detail { /*! \brief Implementation helper for GetArgType */ template ArgType GetArgTypeByIndex(const Call& call, const Op& op, size_t index) { - if (call->args[index]->ty.IsMissing()) { + if (call->args[index]->ty.as().has_value()) { TVM_FFI_VISIT_THROW(InternalError, call) << op << " op should have arguments with defined Type. " << "However, args[" << index << "] has undefined type."; diff --git a/src/relax/script/ir_builder/ir.cc b/src/relax/script/ir_builder/ir.cc index 172a490e5941..df642ee45798 100644 --- a/src/relax/script/ir_builder/ir.cc +++ b/src/relax/script/ir_builder/ir.cc @@ -274,7 +274,7 @@ tvm::Var Emit(const tvm::relax::Expr& expr, const ffi::Optional& anno const tvm::relax::BlockBuilder& block_builder = GetBlockBuilder(); if (annotate_ty.has_value()) { const auto& ty = annotate_ty.value(); - if (expr->ty.IsMissing()) { + if (expr->ty.as().has_value()) { tvm::relax::UpdateType(expr, ty); } else { TVM_FFI_ICHECK(tvm::relax::TypeBaseCheck(ty, GetType(expr)) != diff --git a/src/relax/script/printer/binding.cc b/src/relax/script/printer/binding.cc index 5db6e124de25..4a96a4beab95 100644 --- a/src/relax/script/printer/binding.cc +++ b/src/relax/script/printer/binding.cc @@ -45,7 +45,8 @@ ffi::Optional MatchCastDocTranslate(DocTranslatorObj* d, ffi::AnyView i ->Call({d->Translate(binding->value).value(), d->Translate(binding->ty).value()}); IdDoc lhs = VarDoc(d, binding->var); ffi::Optional annotation = std::nullopt; - if (!binding->var->ty.IsMissing()) annotation = d->Translate(binding->var->ty).value(); + if (!binding->var->ty.as().has_value()) + annotation = d->Translate(binding->var->ty).value(); if (!d->GetExtraConfig("relax.show_all_ty", true)) annotation = std::nullopt; d->Emit(AssignDoc(lhs, rhs, annotation), ffi::GetRef(binding)); return std::nullopt; @@ -112,14 +113,14 @@ ffi::Optional VarBindingDocTranslate(DocTranslatorObj* d, ffi::AnyView } } } - bool inferable = inferred.has_value() && !inferred.value().IsMissing() && + bool inferable = inferred.has_value() && !inferred.value().as().has_value() && ffi::StructuralEqual()(binding->var->ty, inferred.value()); bool explicit_output_type = output_type_argument && inferable; // Primitive aliases need their annotation. Without context-free inference, // keep the binding type and let the parser's deferred inference use it. bool elide_annotation = infer_vdevice || explicit_output_type || (!show_all_ty && !binding->var->ty.as() && inferable); - if (!binding->var->ty.IsMissing() && !elide_annotation) { + if (!binding->var->ty.as().has_value() && !elide_annotation) { annotation = d->Translate(binding->var->ty).value(); } d->Emit(AssignDoc(lhs, rhs.value(), annotation), ffi::GetRef(binding)); @@ -148,7 +149,7 @@ ffi::Optional IfDocTranslate(DocTranslatorObj* d, ffi::AnyView input, Var var = ffi::GetRef(static_cast(destination)); ffi::Optional lhs = VarDoc(d, var); ffi::Optional annotation = std::nullopt; - if (!var->ty.IsMissing()) annotation = d->Translate(var->ty).value(); + if (!var->ty.as().has_value()) annotation = d->Translate(var->ty).value(); ExprDoc condition = d->Translate(branch->cond).value(); d->Emit(IfDoc(condition, RelaxSeqBody(d, branch->true_branch.get(), lhs, annotation, destination), RelaxSeqBody(d, branch->false_branch.get(), lhs, annotation, destination)), diff --git a/src/relax/script/printer/call.cc b/src/relax/script/printer/call.cc index d6ffdd3dde98..a9f1f6a81279 100644 --- a/src/relax/script/printer/call.cc +++ b/src/relax/script/printer/call.cc @@ -62,7 +62,7 @@ namespace { // Relax's named constructors defer result typing until a binding supplies it. // A standalone named call therefore reconstructs only a missing stored result. bool HasRelaxCallResult(const CallNode* call, const ffi::Object* destination) { - if (call->ty.IsMissing()) return true; + if (call->ty.as().has_value()) return true; const auto* var = destination && destination->IsInstance() ? static_cast(destination) : nullptr; @@ -266,8 +266,9 @@ ffi::Optional CallDefaultDocTranslate(DocTranslatorObj* d, ffi::AnyView return d->Translate(call->op).value()->Call(args); } ffi::Optional inferred = std::nullopt; - if (call->op.as() && std::all_of(call->args.begin(), call->args.end(), - [](const Expr& arg) { return !arg->ty.IsMissing(); })) { + if (call->op.as() && std::all_of(call->args.begin(), call->args.end(), [](const Expr& arg) { + return !arg->ty.as().has_value(); + })) { try { inferred = Call::ReinferType(call); } catch (const ffi::Error&) { diff --git a/src/relax/script/printer/function.cc b/src/relax/script/printer/function.cc index 3e7eceb2f2ae..d0ac72a0719a 100644 --- a/src/relax/script/printer/function.cc +++ b/src/relax/script/printer/function.cc @@ -55,14 +55,14 @@ ffi::Optional FunctionDocTranslate(DocTranslatorObj* d, ffi::AnyView in for (const Var& param : func->params) { IdDoc lhs = param_ids[param_index++]; ffi::Optional annotation = std::nullopt; - if (!param->ty.IsMissing()) annotation = d->Translate(param->ty).value(); + if (!param->ty.as().has_value()) annotation = d->Translate(param->ty).value(); AssignDoc argument(lhs, std::nullopt, annotation); d->RecordOrigin(argument, param); args.push_back(argument); } auto signature_candidates = CopyImplicitDefs(d); ffi::Optional ret_type = std::nullopt; - if (!func->ret_ty.IsMissing()) { + if (!func->ret_ty.as().has_value()) { ret_type = d->Translate(func->ret_ty).value(); } ffi::Array decorator_keys; diff --git a/src/relax/transform/decompose_ops.cc b/src/relax/transform/decompose_ops.cc index f0ee758e44d9..66ce48ca92ef 100644 --- a/src/relax/transform/decompose_ops.cc +++ b/src/relax/transform/decompose_ops.cc @@ -144,7 +144,7 @@ Expr DecomposeLayerNorm(const Call& call) { } Expr TensorToShape(const Call& call_node, const BlockBuilder& builder) { - TVM_FFI_ICHECK(!call_node->ty.IsMissing()); + TVM_FFI_ICHECK(!call_node->ty.as().has_value()); Expr expr = call_node->args[0]; const ShapeTypeNode* ty = GetTypeAs(call_node); TVM_FFI_ICHECK(ty); diff --git a/src/relax/transform/normalize.cc b/src/relax/transform/normalize.cc index 0757b4f29005..765b4df94614 100644 --- a/src/relax/transform/normalize.cc +++ b/src/relax/transform/normalize.cc @@ -147,7 +147,7 @@ class NormalizeMutator : public ExprMutatorBase { void VisitBinding_(const VarBindingNode* binding) { Expr new_value = this->VisitExpr(binding->value); - if (binding->var->ty.IsMissing()) { + if (binding->var->ty.as().has_value()) { UpdateType(binding->var, GetType(new_value)); } diff --git a/src/relax/transform/rewrite_dataflow_reshape.cc b/src/relax/transform/rewrite_dataflow_reshape.cc index 90bc41fc2b2a..630e9a538ff7 100644 --- a/src/relax/transform/rewrite_dataflow_reshape.cc +++ b/src/relax/transform/rewrite_dataflow_reshape.cc @@ -125,7 +125,8 @@ class DataflowReshapeRewriter : public ExprMutator { // as the number of elements in the result. There are operators that could have a reshape // pattern that don't meet this requirement (e.g. strided_slice), and they should not be // converted to reshape. - TVM_FFI_ICHECK(!inp->ty.IsMissing() && !call->ty.IsMissing()); + TVM_FFI_ICHECK(!inp->ty.as().has_value() && + !call->ty.as().has_value()); TensorType inp_ty = inp->ty.as_or_throw(); TensorType res_ty = call->ty.as_or_throw(); diff --git a/src/relax/transform/update_vdevice.cc b/src/relax/transform/update_vdevice.cc index 860596c61169..e4650ce676b6 100644 --- a/src/relax/transform/update_vdevice.cc +++ b/src/relax/transform/update_vdevice.cc @@ -43,7 +43,7 @@ class VDeviceMutator : public ExprMutator { Expr VisitExpr(const Expr& expr) final { auto visited_expr = ExprMutator::VisitExpr(expr); - if (!visited_expr->ty.IsMissing()) { + if (!visited_expr->ty.as().has_value()) { auto* tinfo = GetTypeAs(visited_expr); bool unchanged = true; if (tinfo != nullptr) { diff --git a/src/relax/transform/utils.h b/src/relax/transform/utils.h index 3424b4b3645b..328a8f72bfbe 100644 --- a/src/relax/transform/utils.h +++ b/src/relax/transform/utils.h @@ -243,7 +243,8 @@ class SymbolicVarRenewMutator : public ExprMutator { using relax::ExprMutator::VisitExpr_; static Var CopyVar(const VarNode* op, Type ty) { - ffi::Optional ty_annotation = ty.IsMissing() ? std::nullopt : ffi::Optional(ty); + ffi::Optional ty_annotation = + ty.as().has_value() ? std::nullopt : ffi::Optional(ty); if (op->IsInstance()) { return DataflowVar(op->name, std::move(ty_annotation), op->span); } @@ -251,7 +252,7 @@ class SymbolicVarRenewMutator : public ExprMutator { } Type RenewType(const VarNode* op) { - return op->ty.IsMissing() ? op->ty : this->VisitExprDepTypeField(op->ty); + return op->ty.as().has_value() ? op->ty : this->VisitExprDepTypeField(op->ty); } Var RenewVarDefinition(const VarNode* op) { diff --git a/src/script/ir_builder/ir.cc b/src/script/ir_builder/ir.cc index 1a667fb21623..364f51350590 100644 --- a/src/script/ir_builder/ir.cc +++ b/src/script/ir_builder/ir.cc @@ -54,7 +54,7 @@ IRModuleFrame IRModule() { // each dialect registers its own handler that maps a function of that // type to the appropriate ty. inline ffi::Optional GetGlobalVarType(const BaseFunc& func) { - if (!func->ty.IsMissing()) { + if (!func->ty.as().has_value()) { return func->ty; } // Registry: "script.ir_builder.decl_function." — per-function-kind diff --git a/src/script/printer/ir/prim_type.cc b/src/script/printer/ir/prim_type.cc index bf64ce7965bb..b3ef7ccec3b3 100644 --- a/src/script/printer/ir/prim_type.cc +++ b/src/script/printer/ir/prim_type.cc @@ -32,6 +32,16 @@ namespace details { namespace { +ffi::Optional MissingTypeDocTranslate(DocTranslatorObj*, ffi::AnyView, + const ffi::Object*) { + return NamespaceDoc("ir")->Attr("MissingType")->Call({}); +} + +TVM_FFI_STATIC_INIT_BLOCK() { + ffi::reflection::TypeAttrDef().attr( + kDocTranslate, FDocTranslate::FromNative<&MissingTypeDocTranslate>()); +} + ffi::Optional IntImmDocTranslate(DocTranslatorObj* d, ffi::AnyView input, const ffi::Object*) { const auto* imm = diff --git a/src/script/printer/ir/utils.cc b/src/script/printer/ir/utils.cc index 2840ddcb3b52..23769909edc9 100644 --- a/src/script/printer/ir/utils.cc +++ b/src/script/printer/ir/utils.cc @@ -107,8 +107,8 @@ ExprDoc TypeValueImpl(DocTranslatorObj* d, const Type& type, bool dtype_literal) ->Attr("FuncType") ->Call({ListDoc(args), TypeValue(d, function->ret_type, false)}); } - if (type.IsMissing()) { - return NamespaceDoc("ir")->Attr("Type")->Attr("missing")->Call({}); + if (type.as().has_value()) { + return NamespaceDoc("ir")->Attr("MissingType")->Call({}); } if (auto pointer = type.as()) { if (auto primitive = pointer->element_type.as(); diff --git a/src/script/printer/script_printer.cc b/src/script/printer/script_printer.cc index 4c6ce2337c0a..3ad12527c82d 100644 --- a/src/script/printer/script_printer.cc +++ b/src/script/printer/script_printer.cc @@ -63,6 +63,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); + RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); RegisterScriptRepr(); diff --git a/src/sym/z3_prover.cc b/src/sym/z3_prover.cc index 143e698d4764..288de3620145 100644 --- a/src/sym/z3_prover.cc +++ b/src/sym/z3_prover.cc @@ -743,7 +743,7 @@ class Z3Prover::Impl : tvm::ExprFunctor { /// @brief Check if the expression type is supported by z3 integer operations. static bool IsZ3SupportedExpr(const ExprNode* expr) { TVM_FFI_DCHECK(expr != nullptr); - TVM_FFI_DCHECK(!expr->ExprNode::ty.IsMissing()); + TVM_FFI_DCHECK(!expr->ExprNode::ty.as().has_value()); PrimType prim_ty = expr->ExprNode::ty.as_or_throw(); return (prim_ty->dtype.code == static_cast(DLDataTypeCode::kDLInt) || prim_ty->dtype.code == static_cast(DLDataTypeCode::kDLUInt) || diff --git a/src/target/llvm/codegen_llvm.cc b/src/target/llvm/codegen_llvm.cc index f1cd7aa7281c..3b9abaccd599 100644 --- a/src/target/llvm/codegen_llvm.cc +++ b/src/target/llvm/codegen_llvm.cc @@ -2308,7 +2308,7 @@ void CodeGenLLVM::Dispatch_(const BindNode* op) { // Therefore, to have the correct LLVM type for pointers, we may // need to introduce a pointer-cast, even though pointer-to-pointer // casts are not expressible with the `prim::CastNode`. - if (is_pointer && !v->ty.IsMissing()) { + if (is_pointer && !v->ty.as().has_value()) { TVM_FFI_ICHECK(op->value->ty.as()) << "Variable " << op->var << " is a pointer with type " << op->value << ", but is being bound to expression with type " << op->value->ty; diff --git a/src/tirx/ir/buffer_common.h b/src/tirx/ir/buffer_common.h index 26138a7cbd1b..c728e3fad096 100644 --- a/src/tirx/ir/buffer_common.h +++ b/src/tirx/ir/buffer_common.h @@ -41,7 +41,7 @@ namespace tirx { * type. Otherwise the object is nullopt. */ inline std::optional GetPointerType(const Type& type) { - if (!type.IsMissing()) { + if (!type.as().has_value()) { if (auto* ptr_type = type.as()) { if (auto* prim_type = ptr_type->element_type.as()) { return ffi::GetRef(prim_type); diff --git a/src/tirx/ir/function.cc b/src/tirx/ir/function.cc index 4cd2282bc1e1..3a745ac47717 100644 --- a/src/tirx/ir/function.cc +++ b/src/tirx/ir/function.cc @@ -146,7 +146,7 @@ TVM_FFI_INLINE ffi::Expected> PrimFuncMaybeInplaceMut PrimFunc::PrimFunc(ffi::Array params, ffi::Optional body, Type ret_type, DictAttrs attrs, Span span) : BaseFunc(ffi::UnsafeInit{}) { - if (ret_type.IsMissing()) { + if (ret_type.as().has_value()) { ret_type = VoidType(); } diff --git a/src/tirx/op/op.cc b/src/tirx/op/op.cc index 877e5f583ee2..4ca6bedf32dd 100644 --- a/src/tirx/op/op.cc +++ b/src/tirx/op/op.cc @@ -67,7 +67,7 @@ Type GetType(const PrimExpr& expr) { if (auto* ptr = expr.as()) { // If Var has a more refined type annotation, // return the type anotation - if (!ptr->ty.IsMissing()) { + if (!ptr->ty.as().has_value()) { return ptr->ty; } } diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 405729ea0211..f92676e1eae4 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -91,7 +91,7 @@ TVM_FFI_STATIC_INIT_BLOCK() { bool CanTranslateExplicitResultCall(const CallNode* call) { return !call->attrs.defined() && call->ty_args.empty() && call->ty.as() && std::all_of(call->args.begin(), call->args.end(), [](const Expr& arg) { - return !arg->ty.IsMissing() && !arg.as(); + return !arg->ty.as().has_value() && !arg.as(); }); } diff --git a/src/tirx/script/printer/function.cc b/src/tirx/script/printer/function.cc index 4da3b582f871..09c2a5f8d6fd 100644 --- a/src/tirx/script/printer/function.cc +++ b/src/tirx/script/printer/function.cc @@ -134,7 +134,7 @@ void PrintPrimFunc(DocTranslatorObj* d, const tirx::PrimFuncNode* func, ExprDoc ExprStmtDoc(NamespaceDoc("tirx")->Attr("func_attr")->Call({DictDoc(keys, values)}))); } ffi::Optional ret_type = std::nullopt; - if (!func->ret_type.IsMissing() && !IsVoidType(func->ret_type)) { + if (!func->ret_type.as().has_value() && !IsVoidType(func->ret_type)) { ret_type = d->Translate(func->ret_type).value(); } doc = FunctionDoc(IdDoc(name), args, {decorator}, ret_type, body); diff --git a/tests/cpp/expr_test.cc b/tests/cpp/expr_test.cc index 69d692c68238..748f0f4d3704 100644 --- a/tests/cpp/expr_test.cc +++ b/tests/cpp/expr_test.cc @@ -70,8 +70,8 @@ TEST(Expr, RequiredIRReferences) { CheckRequiredIRReference(); CheckRequiredIRReference(); CheckRequiredIRReference(); - EXPECT_TRUE(Type::Missing().IsMissing()); - EXPECT_TRUE(ffi::Any(Type::Missing()).cast().IsMissing()); + EXPECT_TRUE(Type::Missing().as().has_value()); + EXPECT_TRUE(ffi::Any(Type::Missing()).cast().as().has_value()); EXPECT_THROW(ffi::Array({ffi::Any()}).as_or_throw>(), ffi::Error); } diff --git a/tests/python/ir/test_ir_type.py b/tests/python/ir/test_ir_type.py index bda4e5e04e25..7bac117a32d3 100644 --- a/tests/python/ir/test_ir_type.py +++ b/tests/python/ir/test_ir_type.py @@ -33,7 +33,7 @@ def test_missing_type(): missing = tvm.ir.Type.missing() assert isinstance(missing, tvm.ir.Type) - assert missing.is_missing() + assert isinstance(missing, tvm.ir.MissingType) def test_prim_type(): diff --git a/tests/python/ir/test_nonnullable_ir.py b/tests/python/ir/test_nonnullable_ir.py index ebdc3ce725a5..dcbb69ba1103 100644 --- a/tests/python/ir/test_nonnullable_ir.py +++ b/tests/python/ir/test_nonnullable_ir.py @@ -59,7 +59,7 @@ def test_optional_statement_fields_and_roundtrip(): def test_offset_default_and_missing_type_remain_values(): tensor = tvm.tirx.decl_tensor((8,), "float32", elem_offset=None) assert int(tensor.ty.elem_offset) == 0 - assert tvm.ir.Type.missing().is_missing() + assert isinstance(tvm.ir.Type.missing(), tvm.ir.MissingType) # False and zero are valid primitive expressions, not absence. assert int(tvm.tirx.Evaluate(False).value) == 0 assert int(tvm.tirx.Evaluate(0).value) == 0 diff --git a/tests/python/relax/test_analysis_well_formed.py b/tests/python/relax/test_analysis_well_formed.py index 84c190557230..1d925b811448 100644 --- a/tests/python/relax/test_analysis_well_formed.py +++ b/tests/python/relax/test_analysis_well_formed.py @@ -172,7 +172,7 @@ def test_unchecked_call_constructor(): call = tvm.ir.Call.unchecked("relax.add", [x], attrs={"key": 1}, span=span) assert isinstance(call, tvm.ir.Call) assert call.op.same_as(op) - assert call.ty.is_missing() + assert isinstance(call.ty, tvm.ir.MissingType) assert call.span.same_as(span) assert isinstance(call.attrs, tvm.ir.DictAttrs) assert len(call.ty_args) == 0 diff --git a/tests/python/relax/test_expr.py b/tests/python/relax/test_expr.py index 8dbed8951ed1..feb1f3f5704f 100644 --- a/tests/python/relax/test_expr.py +++ b/tests/python/relax/test_expr.py @@ -44,7 +44,7 @@ def _check_json_roundtrip(x): def _check_type_missing(ty): assert isinstance(ty, tvm.ir.Type) - assert ty.is_missing() + assert isinstance(ty, tvm.ir.MissingType) def test_var() -> None: diff --git a/tests/python/script/test_constructor_contracts.py b/tests/python/script/test_constructor_contracts.py index 33dbf12fffc4..53627eca6091 100644 --- a/tests/python/script/test_constructor_contracts.py +++ b/tests/python/script/test_constructor_contracts.py @@ -32,7 +32,7 @@ def test_call_type_and_validation_contract(): for constructor in (ir.Call, I.Call): call = constructor("tirx.exp", [x], span=span) assert isinstance(call, I.Call) - assert call.ty.is_missing() + assert isinstance(call.ty, ir.MissingType) assert call.span.same_as(span) ir.assert_structural_equal( constructor("tirx.exp", [x], ty="float32"),