首先给出TVM中注册自定义函数和调用自定义函数的方法
// 注册函数
TVM_REGISTER_GLOBAL("add").set_body([](TVMArgs args, TVMRetValue* ret) {
int a = args[0];
int b = args[1];
*ret = a + b;
});
// 调用函数
PackedFunc add = runtime::Registry::Get("add");
int result = add(3, 5); // 返回 8
TVM实现注册Lambda函数的set_body函数是一个指向PackedFunc类型的指针.
TVM_REGISTER_GLOBAL的实现如下:
#define REGISTER_GLOBAL(name, func) \
tvm::runtime::FRegistry::Register(name, tvm::runtime::PackedFunc(func))
使用全局哈希表FRegistry注册函数。通过宏REGISTER_GLOBAL("func_name", MyFunction)将函数与名称绑定,后续通过GetPackedFunc("func_name")查找.
使用REGISTER_GLOBAL宏将函数与名称绑定。这个宏会调用FRegistry::Register方法,将函数存储到全局哈希表中。
PackedFunc类型继承自ObjectRef基类,实现了运算符重载,又用make_object函数创建一个PackedFuncSubObj类型对象,这个对象可以储存可调用对象.
PackedFuncSubObj继承自PackedFuncObj, 这是Object的子类,Object实现了引用计数和类型检查,PackedFunObj对函数指针、参数和返回值指针进行了打包。
PackedFuncSubObj类型用std::remove_reference和std::remove_cv进行了类型擦除,对const、volatile和引用进行去壳,移除我们不需要的特性.
PackedFuncSubObj中定义了Extractor提取器结构,提取器内部的Call函数是一个指针,用来调用可调用对象。
接下来解释一下参数和返回值的数据结构。
分别是TVMArgs和TVMRetValue,都使用了联合体TVMValue对数据进行打包并进行了运算符重载和用于数据传递的基本方法。
以上所有的实现基本都在include/tvm/runtime/packed_func.h、include/tvm/runtime/registry.h和src/runtime/registry.cc
Python封装了ctypes库,能够通过name查找全局注册的C++函数并获得函数句柄,调用后得到传回的返回值。
其中对数据类型的包装也是TVM实现任意语言互相调用的关键.