Skip to content

Commit 3f75a27

Browse files
authored
Update FromNative call sites for latest TVM FFI (#20483)
- Advance the bundled TVM FFI submodule to the native hook convention while retaining the published package requirement. - Model typed hooks by their successful return type and use `CallExpected()` where TVM handles errors explicitly. - Register throwing Relax validators and type inference hooks directly with `FromNative`, removing redundant exception adapters.
1 parent a7fffaa commit 3f75a27

27 files changed

Lines changed: 111 additions & 466 deletions

‎include/tvm/ir/op.h‎

Lines changed: 13 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ template <typename>
4444
class OpAttrMap;
4545

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

4949
/*! \brief An operator argument's name and documentation. */
5050
class ArgumentInfoNode : public ffi::Object {
@@ -171,9 +171,11 @@ class Op : public Expr {
171171
if (TVM_FFI_PREDICT_FALSE(!call)) {
172172
ThrowInvalidCall(get());
173173
}
174-
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
174+
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
175175
// set_validator stores only an owning NativeFunction with this signature.
176-
ffi::details::AnyUnsafe::CopyFromAnyViewAfterCheck<View>(validator)(call).value();
176+
ffi::details::AnyUnsafe::CopyFromAnyViewAfterCheck<View>(validator)
177+
.CallExpected(call)
178+
.value();
177179
}
178180
}
179181
TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(Op, Expr, OpNode);
@@ -502,7 +504,7 @@ class OpDef {
502504
auto updated = ffi::make_object<OpNode>();
503505
(ApplySignatureTrait(updated.get(), specs), ...);
504506
if (op_->validator == nullptr) {
505-
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
507+
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
506508
set_validator(View::FromNative<&ValidateSignature<Specs...>>());
507509
get()->validator_is_custom = false;
508510
}
@@ -531,10 +533,10 @@ class OpDef {
531533
/*!
532534
* \brief Register a validator that checks a Call with this operator.
533535
*
534-
* The callback accepts a `const CallNode*` and returns Expected<void>, with
535-
* an error for invalid input. A native function pointer can be bound with
536-
* NativeFunctionView::FromNative; a borrowed packed function may also be
537-
* passed while it remains alive for this call. The setter retains an owning
536+
* The callback accepts a `const CallNode*` and reports invalid input by
537+
* throwing or returning an `Expected<void>` error. A native function pointer
538+
* can be bound with NativeFunctionView::FromNative. A borrowed packed function
539+
* may also be passed while it remains alive for this call. The setter retains an owning
538540
* copy, so the original packed function may then be destroyed.
539541
* Ordinary Call construction invokes it; Call::Unchecked skips that initial
540542
* check, while Relax normalization and well-formedness may validate later.
@@ -547,13 +549,11 @@ class OpDef {
547549
* \param override Whether to replace the current validator.
548550
* \return This builder.
549551
*/
550-
OpDef& set_validator(
551-
ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)> validator,
552-
bool override = false) {
552+
OpDef& set_validator(ffi::reflection::NativeFunctionView<void(const CallNode*)> validator,
553+
bool override = false) {
553554
TVM_FFI_CHECK(override || op_->validator == nullptr, ValueError)
554555
<< "Validator of " << op_->name << " is already registered";
555-
get()->validator =
556-
ffi::reflection::NativeFunction<ffi::Expected<void>(const CallNode*)>::From(validator);
556+
get()->validator = ffi::reflection::NativeFunction<void(const CallNode*)>::From(validator);
557557
get()->validator_is_custom = true;
558558
return *this;
559559
}

‎src/ir/expr.cc‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1213,7 +1213,7 @@ Type Call::ReinferType(const CallNode* call) {
12131213
static auto infer_type = Op::GetAttrMap<FInferType>("FInferType");
12141214
TVM_FFI_CHECK(infer_type.count(op.value()), ValueError)
12151215
<< "No context-free FInferType hook is registered for " << op.value();
1216-
Type result = infer_type[op.value()](call).value();
1216+
Type result = infer_type[op.value()].CallExpected(call).value();
12171217
TVM_FFI_CHECK(!result.IsMissing(), InternalError)
12181218
<< "FInferType for " << op.value() << " returned Type::Missing()";
12191219
return result;

‎src/ir/op.cc‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -208,8 +208,8 @@ void SetOpSignature(Op op, const ffi::Array<ffi::String>& arg_names,
208208
auto var_ty_args_info = MakeTailInfo(var_ty_args);
209209
auto* node = const_cast<OpNode*>(op.operator->());
210210
if (!node->validator_is_custom) {
211-
using View = ffi::reflection::NativeFunctionView<ffi::Expected<void>(const CallNode*)>;
212-
node->validator = ffi::reflection::NativeFunction<ffi::Expected<void>(const CallNode*)>::From(
211+
using View = ffi::reflection::NativeFunctionView<void(const CallNode*)>;
212+
node->validator = ffi::reflection::NativeFunction<void(const CallNode*)>::From(
213213
View::FromNative<&ValidateCountSignature>());
214214
}
215215
node->args_info = std::move(args_info);

‎src/relax/op/ccl/ccl.cc‎

Lines changed: 3 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -51,14 +51,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
5151
refl::GlobalDef().def("relax.op.ccl.allreduce", allreduce);
5252
}
5353

54-
ffi::Expected<Type> InferTypeAllReduce(const CallNode* call_node) noexcept try {
54+
Type InferTypeAllReduce(const CallNode* call_node) {
5555
const Call call = ffi::GetRef<Call>(call_node);
5656
TensorType input_ty = GetUnaryInputTensorType(call);
5757
return input_ty;
58-
} catch (const ffi::Error& error) {
59-
return ffi::Unexpected(error);
60-
} catch (const std::exception& error) {
61-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
6258
}
6359

6460
TVM_FFI_STATIC_INIT_BLOCK() {
@@ -86,7 +82,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
8682
refl::GlobalDef().def("relax.op.ccl.allgather", allgather);
8783
}
8884

89-
ffi::Expected<Type> InferTypeAllGather(const CallNode* call_node) noexcept try {
85+
Type InferTypeAllGather(const CallNode* call_node) {
9086
const Call call = ffi::GetRef<Call>(call_node);
9187
TensorType input_ty = GetUnaryInputTensorType(call);
9288

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

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

130-
ffi::Expected<Type> InferTypeBroadcastFromZero(const CallNode* call_node) noexcept try {
122+
Type InferTypeBroadcastFromZero(const CallNode* call_node) {
131123
const Call call = ffi::GetRef<Call>(call_node);
132124
TensorType input_ty = GetUnaryInputTensorType(call);
133125
return input_ty;
134-
} catch (const ffi::Error& error) {
135-
return ffi::Unexpected(error);
136-
} catch (const std::exception& error) {
137-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
138126
}
139127

140128
TVM_FFI_STATIC_INIT_BLOCK() {

‎src/relax/op/image/resize.cc‎

Lines changed: 4 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -64,7 +64,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
6464
refl::GlobalDef().def("relax.op.image.resize2d", resize2d);
6565
}
6666

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

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

122118
InferLayoutOutput InferLayoutResize2d(
@@ -185,7 +181,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
185181
refl::GlobalDef().def("relax.op.image.resize3d", resize3d);
186182
}
187183

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

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

244236
InferLayoutOutput InferLayoutResize3d(
@@ -299,7 +291,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
299291
refl::GlobalDef().def("relax.op.image.grid_sample", grid_sample);
300292
}
301293

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

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

367355
TVM_FFI_STATIC_INIT_BLOCK() {
@@ -390,7 +378,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
390378
refl::GlobalDef().def("relax.op.image.affine_grid", affine_grid);
391379
}
392380

393-
ffi::Expected<Type> InferTypeAffineGrid(const CallNode* call_node) noexcept try {
381+
Type InferTypeAffineGrid(const CallNode* call_node) {
394382
const Call call = ffi::GetRef<Call>(call_node);
395383
if (call->args.size() != 2) {
396384
TVM_FFI_VISIT_THROW(ValueError, call)
@@ -463,10 +451,6 @@ ffi::Expected<Type> InferTypeAffineGrid(const CallNode* call_node) noexcept try
463451
}
464452

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

472456
TVM_FFI_STATIC_INIT_BLOCK() {

‎src/relax/op/memory/view.cc‎

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -415,18 +415,14 @@ TVM_FFI_STATIC_INIT_BLOCK() {
415415
refl::GlobalDef().def("relax.op.memory.ensure_zero_offset", ensure_zero_offset);
416416
}
417417

418-
ffi::Expected<Type> InferTypeEnsureZeroOffset(const CallNode* call_node) noexcept try {
418+
Type InferTypeEnsureZeroOffset(const CallNode* call_node) {
419419
const Call call = ffi::GetRef<Call>(call_node);
420420
if (call->args.size() != 1) {
421421
TVM_FFI_VISIT_THROW(ValueError, call)
422422
<< "Operator " << call->op << " should receive 1 argument, "
423423
<< "but received " << call->args;
424424
}
425425
return GetType(call->args[0]);
426-
} catch (const ffi::Error& error) {
427-
return ffi::Unexpected(error);
428-
} catch (const std::exception& error) {
429-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
430426
}
431427

432428
Expr LowerBuiltinEnsureZeroOffset(const BlockBuilder& bb, const Call& call) {

‎src/relax/op/nn/nn.cc‎

Lines changed: 6 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -204,7 +204,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
204204
refl::GlobalDef().def("relax.op.nn.prelu", prelu);
205205
}
206206

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

222222
return data_ty;
223-
} catch (const ffi::Error& error) {
224-
return ffi::Unexpected(error);
225-
} catch (const std::exception& error) {
226-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
227223
}
228224

229225
InferLayoutOutput InferLayoutPRelu(
@@ -275,7 +271,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
275271
refl::GlobalDef().def("relax.op.nn.softmax", softmax);
276272
}
277273

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

296292
return data_ty;
297-
} catch (const ffi::Error& error) {
298-
return ffi::Unexpected(error);
299-
} catch (const std::exception& error) {
300-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
301293
}
302294

303295
InferLayoutOutput InferLayoutSoftmax(
@@ -365,7 +357,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
365357
refl::GlobalDef().def("relax.op.nn.pad", pad);
366358
}
367359

368-
ffi::Expected<Type> InferTypePad(const CallNode* call_node) noexcept try {
360+
Type InferTypePad(const CallNode* call_node) {
369361
const Call call = ffi::GetRef<Call>(call_node);
370362
ffi::Array<TensorType> input_ty = GetInputTensorType(call);
371363
const auto* attrs = call->attrs.as<PadAttrs>();
@@ -388,10 +380,6 @@ ffi::Expected<Type> InferTypePad(const CallNode* call_node) noexcept try {
388380
return TensorType(input_ty[0]->dtype, ndim);
389381
}
390382
return TensorType(ShapeExpr(out_shape), input_ty[0]->dtype);
391-
} catch (const ffi::Error& error) {
392-
return ffi::Unexpected(error);
393-
} catch (const std::exception& error) {
394-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
395383
}
396384

397385
TVM_FFI_STATIC_INIT_BLOCK() {
@@ -415,7 +403,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
415403
refl::GlobalDef().def("relax.op.nn.pixel_shuffle", pixel_shuffle);
416404
}
417405

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

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

474458
TVM_FFI_STATIC_INIT_BLOCK() {
@@ -989,14 +973,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
989973
refl::GlobalDef().def("relax.op.nn.dropout", dropout);
990974
}
991975

992-
ffi::Expected<Type> InferTypeDropout(const CallNode* call_node) noexcept try {
976+
Type InferTypeDropout(const CallNode* call_node) {
993977
const Call call = ffi::GetRef<Call>(call_node);
994978
TensorType data_ty = GetUnaryInputTensorType(call);
995979
return TupleType({data_ty, data_ty});
996-
} catch (const ffi::Error& error) {
997-
return ffi::Unexpected(error);
998-
} catch (const std::exception& error) {
999-
return ffi::Unexpected(ffi::Error("InternalError", error.what(), ""));
1000980
}
1001981

1002982
TVM_FFI_STATIC_INIT_BLOCK() {
@@ -1314,7 +1294,7 @@ TVM_FFI_STATIC_INIT_BLOCK() {
13141294
refl::GlobalDef().def("relax.op.nn.batch_flatten", batch_flatten);
13151295
}
13161296

1317-
ffi::Expected<Type> InferTypeBatchFlatten(const CallNode* call_node) noexcept try {
1297+
Type InferTypeBatchFlatten(const CallNode* call_node) {
13181298
const Call call = ffi::GetRef<Call>(call_node);
13191299
TensorType data_ty = GetUnaryInputTensorType(call);
13201300

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

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

13531329
TVM_FFI_STATIC_INIT_BLOCK() {

0 commit comments

Comments
 (0)