Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/reference/api/python/script/ir_builder.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
***************************************
Expand Down
35 changes: 28 additions & 7 deletions include/tvm/ir/base_expr.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<MissingTypeNode>();
}

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.
*
Expand Down Expand Up @@ -311,18 +332,18 @@ 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;
// span does not participate in structural equal and hash.
refl::ObjectDef<ExprNode>()
.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;
Expand Down
2 changes: 1 addition & 1 deletion include/tvm/relax/expr_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -524,7 +524,7 @@ class ExprMutatorBase : public ExprFunctor<Expr(const Expr&)> {
bool VisitAndCheckTypeFieldUnchanged(const ffi::ObjectRef& ty) {
if (const TypeNode* ty_node = ty.as<TypeNode>()) {
Type type = ffi::GetRef<Type>(ty_node);
return type.IsMissing() || this->VisitExprDepTypeField(type).same_as(ty);
return type.as<MissingType>().has_value() || this->VisitExprDepTypeField(type).same_as(ty);
} else {
return true;
}
Expand Down
6 changes: 3 additions & 3 deletions include/tvm/relax/type.h
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ class TensorTypeNode : public TypeNode {
ffi::Optional<ffi::Array<PrimExpr>> 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<MissingType>().has_value()) return {};
if (const auto* shape_ty = shape_expr->ty.as<ShapeTypeNode>()) {
return shape_ty->values;
}
Expand Down Expand Up @@ -366,7 +366,7 @@ inline ffi::Optional<T> MatchType(const Expr& expr) {
*/
template <typename T>
inline const T* GetTypeAs(const Expr& expr) {
TVM_FFI_ICHECK(!expr->ty.IsMissing())
TVM_FFI_ICHECK(!expr->ty.as<MissingType>().has_value())
<< "The type is not populated, check if you have normalized the expr";
return expr->ty.as<T>();
}
Expand All @@ -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<MissingType>().has_value())
<< "The type is not populated, check if you have normalized the expr";
return expr->ty;
}
Expand Down
3 changes: 1 addition & 2 deletions include/tvm/relax/type_functor.h
Original file line number Diff line number Diff line change
Expand Up @@ -73,8 +73,7 @@ class TypeFunctor<R(const Type& n, Args...)> {
*/
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<MissingType>().has_value()) << "TypeFunctor cannot visit Type::Missing()";
static FType vtable = InitVTable();
return vtable(n, this, std::forward<Args>(args)...);
}
Expand Down
2 changes: 1 addition & 1 deletion include/tvm/tirx/op.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<MissingType>().has_value()) return false;
if (const auto* ptr_type = type.as<PointerTypeNode>()) {
if (const auto* prim_type = ptr_type->element_type.as<PrimTypeNode>()) {
return prim_type->dtype == element_type;
Expand Down
12 changes: 11 additions & 1 deletion python/tvm/ir/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/ir/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
20 changes: 14 additions & 6 deletions python/tvm/ir/type.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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."""
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/relax/expr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
4 changes: 2 additions & 2 deletions python/tvm/relax/frontend/nn/subroutine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
5 changes: 2 additions & 3 deletions python/tvm/relax/op/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/relax/testing/ast_printer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
2 changes: 2 additions & 0 deletions python/tvm/script/ir_builder/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
DataTypeImm,
FuncType,
GenericConst,
MissingType,
PrimType,
Range,
StringImm,
Expand Down Expand Up @@ -67,6 +68,7 @@
"GenericConst",
"IRBuilder",
"IRModuleFrame",
"MissingType",
"PrimType",
"Range",
"StringImm",
Expand Down
7 changes: 4 additions & 3 deletions src/ir/expr.cc
Original file line number Diff line number Diff line change
Expand Up @@ -725,7 +725,7 @@ Tuple::Tuple(ffi::Array<Expr> fields, Span span) : Expr(ffi::UnsafeInit{}) {
ffi::Optional<Type> tuple_ty = [&]() -> ffi::Optional<Type> {
ffi::Array<Type> field_ty;
for (const Expr& field : fields) {
if (field->ty.IsMissing()) {
if (field->ty.as<MissingType>().has_value()) {
return std::nullopt;
}
field_ty.push_back(field->ty);
Expand Down Expand Up @@ -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<MissingType>().has_value(), TypeError)
<< "GenericConst requires an expression type";
TVM_FFI_CHECK(!value.as<ffi::BigInt>() && !value.as<bool>() && !value.as<double>() &&
!value.as<ffi::String>(),
TypeError)
Expand Down Expand Up @@ -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<MissingType>().has_value(), InternalError)
<< "FInferType for " << op.value() << " returned Type::Missing()";
return result;
}
Expand Down
2 changes: 1 addition & 1 deletion src/ir/prim/const_fold.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<MissingType>().has_value());
const auto* prim_ty = expr->ExprNode::ty.as<PrimTypeNode>();
TVM_FFI_DCHECK(prim_ty != nullptr);
return IsIndexType(prim_ty->dtype);
Expand Down
2 changes: 1 addition & 1 deletion src/ir/prim/expr.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<MissingType>().has_value());
const auto* prim_ty = node->ExprNode::ty.as<PrimTypeNode>();
TVM_FFI_DCHECK(prim_ty != nullptr);
return prim_ty;
Expand Down
2 changes: 1 addition & 1 deletion src/ir/prim/op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<MissingType>().has_value());
const auto* prim_ty = node->ExprNode::ty.as<PrimTypeNode>();
TVM_FFI_DCHECK(prim_ty != nullptr);
return prim_ty;
Expand Down
34 changes: 20 additions & 14 deletions src/ir/type.cc
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@ ffi::ObjectPtr<PrimTypeNode> GetCachedPrimTypeNode(DLDataType dtype) {

TVM_FFI_INLINE ffi::Expected<ffi::Optional<ffi::VisitInterrupt>> 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;
}

Expand Down Expand Up @@ -246,14 +246,7 @@ TVM_FFI_INLINE ffi::Expected<ffi::UnchangedOr<ffi::Any>> TupleTypeMaybeInplaceMu

} // namespace

Type Type::Missing() {
static Type missing = []() {
Type type(ffi::UnsafeInit{});
type.data_ = ffi::make_object<TypeNode>();
return type;
}();
return missing;
}
Type Type::Missing() { return MissingType(); }

TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
Expand All @@ -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<MissingTypeNode>();
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<MissingTypeNode>()
.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<AnyTypeNode> n = ffi::make_object<AnyTypeNode>();
Expand Down Expand Up @@ -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<MissingType>().has_value())
<< "PointerType element_type cannot be Type::Missing()";
ffi::ObjectPtr<PointerTypeNode> n = ffi::make_object<PointerTypeNode>();
if (storage_scope.empty()) {
n->storage_scope = "global";
Expand Down
Loading
Loading