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
26 changes: 13 additions & 13 deletions include/tvm/ir/op.h
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ template <typename>
class OpAttrMap;

/*! \brief Infer a Call's result type from its explicit inputs without builder state. */
using FInferType = ffi::reflection::NativeFunctionView<ffi::Expected<Type>(const CallNode* call)>;
using FInferType = ffi::reflection::NativeFunctionView<Type(const CallNode* call)>;

/*! \brief An operator argument's name and documentation. */
class ArgumentInfoNode : public ffi::Object {
Expand Down Expand Up @@ -171,9 +171,11 @@ class Op : public Expr {
if (TVM_FFI_PREDICT_FALSE(!call)) {
ThrowInvalidCall(get());
}
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
// set_validator stores only an owning NativeFunction with this signature.
ffi::details::AnyUnsafe::CopyFromAnyViewAfterCheck<View>(validator)(call).value();
ffi::details::AnyUnsafe::CopyFromAnyViewAfterCheck<View>(validator)
.CallExpected(call)
.value();
}
}
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Op, Expr, OpNode);
Expand Down Expand Up @@ -502,7 +504,7 @@ class OpDef {
auto updated = ffi::make_object<OpNode>();
(ApplySignatureTrait(updated.get(), specs), ...);
if (op_->validator == nullptr) {
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
set_validator(View::FromNative<&ValidateSignature<Specs...>>());
get()->validator_is_custom = false;
}
Expand Down Expand Up @@ -531,10 +533,10 @@ class OpDef {
/*!
* \brief Register a validator that checks a Call with this operator.
*
* The callback accepts a `const CallNode*` and returns Expected<void>, with
* an error for invalid input. A native function pointer can be bound with
* NativeFunctionView::FromNative; a borrowed packed function may also be
* passed while it remains alive for this call. The setter retains an owning
* The callback accepts a `const CallNode*` and reports invalid input by
* throwing or returning an `Expected<void>` error. A native function pointer
* can be bound with NativeFunctionView::FromNative. A borrowed packed function
* may also be passed while it remains alive for this call. The setter retains an owning
* copy, so the original packed function may then be destroyed.
* Ordinary Call construction invokes it; Call::Unchecked skips that initial
* check, while Relax normalization and well-formedness may validate later.
Expand All @@ -547,13 +549,11 @@ class OpDef {
* \param override Whether to replace the current validator.
* \return This builder.
*/
OpDef& set_validator(
ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)> validator,
bool override = false) {
OpDef& set_validator(ffi::reflection::NativeFunctionView<void(const CallNode*)> validator,
bool override = false) {
TVM_FFI_CHECK(override || op_->validator == nullptr, ValueError)
<< "Validator of " << op_->name << " is already registered";
get()->validator =
ffi::reflection::NativeFunction<ffi::Expected<void>(const CallNode*)>::From(validator);
get()->validator = ffi::reflection::NativeFunction<void(const CallNode*)>::From(validator);
get()->validator_is_custom = true;
return *this;
}
Expand Down
2 changes: 1 addition & 1 deletion src/ir/expr.cc
Original file line number Diff line number Diff line change
Expand Up @@ -1213,7 +1213,7 @@ Type Call::ReinferType(const CallNode* call) {
static auto infer_type = Op::GetAttrMap<FInferType>("FInferType");
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()](call).value();
Type result = infer_type[op.value()].CallExpected(call).value();
TVM_FFI_CHECK(!result.IsMissing(), InternalError)
<< "FInferType for " << op.value() << " returned Type::Missing()";
return result;
Expand Down
4 changes: 2 additions & 2 deletions src/ir/op.cc
Original file line number Diff line number Diff line change
Expand Up @@ -208,8 +208,8 @@ void SetOpSignature(Op op, const ffi::Array<ffi::String>& arg_names,
auto var_ty_args_info = MakeTailInfo(var_ty_args);
auto* node = const_cast<OpNode*>(op.operator->());
if (!node->validator_is_custom) {
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
node->validator = ffi::reflection::NativeFunction<ffi::Expected<void>(const CallNode*)>::From(
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
node->validator = ffi::reflection::NativeFunction<void(const CallNode*)>::From(
View::FromNative<&ValidateCountSignature>());
}
node->args_info = std::move(args_info);
Expand Down
18 changes: 3 additions & 15 deletions src/relax/op/ccl/ccl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -51,14 +51,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.ccl.allreduce", allreduce);
}

ffi::Expected<Type> InferTypeAllReduce(const CallNode* call_node) noexcept try {
Type InferTypeAllReduce(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType input_ty = GetUnaryInputTensorType(call);
return input_ty;
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down Expand Up @@ -86,7 +82,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.ccl.allgather", allgather);
}

ffi::Expected<Type> InferTypeAllGather(const CallNode* call_node) noexcept try {
Type InferTypeAllGather(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType input_ty = GetUnaryInputTensorType(call);

Expand All @@ -101,10 +97,6 @@ ffi::Expected<Type> InferTypeAllGather(const CallNode* call_node) noexcept try {
ffi::Array<PrimExpr> output_shape = input_shape.value();
output_shape.Set(0, floor(output_shape[0] * num_workers));
return TensorType(ShapeExpr(output_shape), output_dtype, input_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand All @@ -127,14 +119,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.ccl.broadcast_from_worker0", broadcast_from_worker0);
}

ffi::Expected<Type> InferTypeBroadcastFromZero(const CallNode* call_node) noexcept try {
Type InferTypeBroadcastFromZero(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType input_ty = GetUnaryInputTensorType(call);
return input_ty;
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down
24 changes: 4 additions & 20 deletions src/relax/op/image/resize.cc
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.image.resize2d", resize2d);
}

ffi::Expected<Type> InferTypeResize2D(const CallNode* call_node) noexcept try {
Type InferTypeResize2D(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
if (call->args.size() != 2) {
TVM_FFI_VISIT_THROW(ValueError, call)
Expand Down Expand Up @@ -113,10 +113,6 @@ ffi::Expected<Type> InferTypeResize2D(const CallNode* call_node) noexcept try {

ffi::Array<PrimExpr> out_shape = data2NCHW.BackwardShape(out_NCHW_shape);
return TensorType(ShapeExpr(out_shape), out_dtype, data_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

InferLayoutOutput InferLayoutResize2d(
Expand Down Expand Up @@ -185,7 +181,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.image.resize3d", resize3d);
}

ffi::Expected<Type> InferTypeResize3D(const CallNode* call_node) noexcept try {
Type InferTypeResize3D(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
if (call->args.size() != 2) {
TVM_FFI_VISIT_THROW(ValueError, call)
Expand Down Expand Up @@ -235,10 +231,6 @@ ffi::Expected<Type> InferTypeResize3D(const CallNode* call_node) noexcept try {

ffi::Array<PrimExpr> out_shape = data2NCDHW.BackwardShape(out_NCDHW_shape);
return TensorType(ShapeExpr(out_shape), out_dtype, data_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

InferLayoutOutput InferLayoutResize3d(
Expand Down Expand Up @@ -299,7 +291,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.image.grid_sample", grid_sample);
}

ffi::Expected<Type> InferTypeGridSample(const CallNode* call_node) noexcept try {
Type InferTypeGridSample(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
if (call->args.size() != 2) {
TVM_FFI_VISIT_THROW(ValueError, call)
Expand Down Expand Up @@ -358,10 +350,6 @@ ffi::Expected<Type> InferTypeGridSample(const CallNode* call_node) noexcept try

ffi::Array<PrimExpr> out_shape = data2tgt.BackwardShape(out_tgt_shape);
return TensorType(ShapeExpr(out_shape), out_dtype, data_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down Expand Up @@ -390,7 +378,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.image.affine_grid", affine_grid);
}

ffi::Expected<Type> InferTypeAffineGrid(const CallNode* call_node) noexcept try {
Type InferTypeAffineGrid(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
if (call->args.size() != 2) {
TVM_FFI_VISIT_THROW(ValueError, call)
Expand Down Expand Up @@ -463,10 +451,6 @@ ffi::Expected<Type> InferTypeAffineGrid(const CallNode* call_node) noexcept try
}

return TensorType(ShapeExpr(out_shape), out_dtype, data_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down
6 changes: 1 addition & 5 deletions src/relax/op/memory/view.cc
Original file line number Diff line number Diff line change
Expand Up @@ -415,18 +415,14 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.memory.ensure_zero_offset", ensure_zero_offset);
}

ffi::Expected<Type> InferTypeEnsureZeroOffset(const CallNode* call_node) noexcept try {
Type InferTypeEnsureZeroOffset(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
if (call->args.size() != 1) {
TVM_FFI_VISIT_THROW(ValueError, call)
<< "Operator " << call->op << " should receive 1 argument, "
<< "but received " << call->args;
}
return GetType(call->args[0]);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

Expr LowerBuiltinEnsureZeroOffset(const BlockBuilder& bb, const Call& call) {
Expand Down
36 changes: 6 additions & 30 deletions src/relax/op/nn/nn.cc
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.prelu", prelu);
}

ffi::Expected<Type> InferTypePRelu(const CallNode* call_node) noexcept try {
Type InferTypePRelu(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType data_ty = GetUnaryInputTensorType(call);
if (data_ty->IsUnknownNdim()) {
Expand All @@ -220,10 +220,6 @@ ffi::Expected<Type> InferTypePRelu(const CallNode* call_node) noexcept try {
NormalizeAxis(call, data_ty->ndim, attrs->axis);

return data_ty;
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

InferLayoutOutput InferLayoutPRelu(
Expand Down Expand Up @@ -275,7 +271,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.softmax", softmax);
}

ffi::Expected<Type> InferTypeSoftmax(const CallNode* call_node) noexcept try {
Type InferTypeSoftmax(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType data_ty = GetUnaryInputTensorType(call);
if (data_ty->IsUnknownNdim()) {
Expand All @@ -294,10 +290,6 @@ ffi::Expected<Type> InferTypeSoftmax(const CallNode* call_node) noexcept try {
NormalizeAxis(call, data_ty->ndim, attrs->axis);

return data_ty;
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

InferLayoutOutput InferLayoutSoftmax(
Expand Down Expand Up @@ -365,7 +357,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.pad", pad);
}

ffi::Expected<Type> InferTypePad(const CallNode* call_node) noexcept try {
Type InferTypePad(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
ffi::Array<TensorType> input_ty = GetInputTensorType(call);
const auto* attrs = call->attrs.as<PadAttrs>();
Expand All @@ -388,10 +380,6 @@ ffi::Expected<Type> InferTypePad(const CallNode* call_node) noexcept try {
return TensorType(input_ty[0]->dtype, ndim);
}
return TensorType(ShapeExpr(out_shape), input_ty[0]->dtype);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand All @@ -415,7 +403,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.pixel_shuffle", pixel_shuffle);
}

ffi::Expected<Type> InferTypePixelShuffle(const CallNode* call_node) noexcept try {
Type InferTypePixelShuffle(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
ffi::Array<TensorType> input_ty = GetInputTensorType(call);
const auto* attrs = call->attrs.as<PixelShuffleAttrs>();
Expand Down Expand Up @@ -465,10 +453,6 @@ ffi::Expected<Type> InferTypePixelShuffle(const CallNode* call_node) noexcept tr
}

return TensorType(ShapeExpr(out_shape), input->dtype);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down Expand Up @@ -989,14 +973,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.dropout", dropout);
}

ffi::Expected<Type> InferTypeDropout(const CallNode* call_node) noexcept try {
Type InferTypeDropout(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType data_ty = GetUnaryInputTensorType(call);
return TupleType({data_ty, data_ty});
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down Expand Up @@ -1314,7 +1294,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
refl::GlobalDef().def("relax.op.nn.batch_flatten", batch_flatten);
}

ffi::Expected<Type> InferTypeBatchFlatten(const CallNode* call_node) noexcept try {
Type InferTypeBatchFlatten(const CallNode* call_node) {
const Call call = ffi::GetRef<Call>(call_node);
TensorType data_ty = GetUnaryInputTensorType(call);

Expand Down Expand Up @@ -1344,10 +1324,6 @@ ffi::Expected<Type> InferTypeBatchFlatten(const CallNode* call_node) noexcept tr
}

return TensorType(ShapeExpr({batch_dim, flat_dim}), data_ty->dtype, data_ty->vdevice);
} catch (const ffi::Error& error) {
return ffi::Unexpected(error);
} catch (const std::exception& error) {
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
}

TVM_FFI_STATIC_INIT_BLOCK() {
Expand Down
Loading
Loading