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
27 changes: 26 additions & 1 deletion include/tvm/runtime/crt/func_registry.h
Original file line number Diff line number Diff line change
Expand Up @@ -42,14 +42,39 @@ typedef struct TVMFuncRegistry {
/*! \brief Names of registered functions, concatenated together and separated by \0.
* An additional \0 is present at the end of the concatenated blob to mark the end.
*
* Byte 0 is the number of functions in `funcs`.
* Byte 0 and 1 are the number of functions in `funcs`.
*/
const char* names;

/*! \brief Function pointers, in the same order as their names in `names`. */
const TVMBackendPackedCFunc* funcs;
} TVMFuncRegistry;

/*!
* \brief Get the of the number of functions from registry.
*
* \param reg TVMFunctionRegistry instance that contains the function.
* \return The number of functions from registry.
*/
uint16_t TVMFuncRegistry_GetNumFuncs(const TVMFuncRegistry* reg);

/*!
* \brief Set the number of functions to registry.
*
* \param reg TVMFunctionRegistry instance that contains the function.
* \param num_funcs The number of functions
* \return 0 when successful.
*/
int TVMFuncRegistry_SetNumFuncs(const TVMFuncRegistry* reg, const uint16_t num_funcs);

/*!
* \brief Get the address of 0th function from registry.
*
* \param reg TVMFunctionRegistry instance that contains the function.
* \return the address of 0th function from registry
*/
const char* TVMFuncRegistry_Get0thFunctionName(const TVMFuncRegistry* reg);

/*!
* \brief Get packed function from registry by name.
*
Expand Down
79 changes: 0 additions & 79 deletions python/tvm/micro/func_registry.py

This file was deleted.

2 changes: 1 addition & 1 deletion src/runtime/crt/aot_executor_module/aot_executor_module.c
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,7 @@ static const TVMBackendPackedCFunc aot_executor_registry_funcs[] = {
};

static const TVMFuncRegistry aot_executor_registry = {
"\x0aget_input\0"
"\x0a\0get_input\0"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you need to update src/runtime/crt/graph_executor_module/graph_executor_module.c too?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sorry, I think that I forget to update src/runtime/crt/graph_executor_module/graph_executor_module.c in #10014.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@A1245967 @areusch There's also python/tvm/micro/func_registry.py which seems to need updating too, although it is unclear if this function graph_json_to_c_func_registry() is used anywhere.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

good catch, done

@A1245967 A1245967 May 19, 2022

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think this line in python/tvm/micro/func_registry.py can be replaced with the following lines.

 import struct
 encoded_NumFuncs = [f"\\{v:03o}" for v in struct.pack("H", len(funcs))]
 encoded_funcs = "".join(encoded_NumFuncs) + "\\0".join(funcs)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@A1245967 that's a good point. I actually looked at this further and I don't see this function called. I think this approach was from before we emitted the FuncRegistry in codegen. I just went ahead and deleted that file.

"get_input_index\0"
"get_input_info\0"
"get_num_inputs\0"
Expand Down
39 changes: 29 additions & 10 deletions src/runtime/crt/common/func_registry.c
Original file line number Diff line number Diff line change
Expand Up @@ -60,14 +60,29 @@ int strcmp_cursor(const char** cursor, const char* name) {
return return_value;
}

uint16_t TVMFuncRegistry_GetNumFuncs(const TVMFuncRegistry* reg) {
uint16_t num_funcs;
memcpy(&num_funcs, reg->names, sizeof(num_funcs));
return num_funcs;
}

int TVMFuncRegistry_SetNumFuncs(const TVMFuncRegistry* reg, const uint16_t num_funcs) {
memcpy((char*)reg->names, &num_funcs, sizeof(num_funcs));
return 0;
}

const char* TVMFuncRegistry_Get0thFunctionName(const TVMFuncRegistry* reg) {
// NOTE: first function name starts at index 2 to skip num_funcs.
return (reg->names + sizeof(uint16_t));
}

tvm_crt_error_t TVMFuncRegistry_Lookup(const TVMFuncRegistry* reg, const char* name,
tvm_function_index_t* function_index) {
tvm_function_index_t idx;
const char* reg_name_ptr;
const char* reg_name_ptr = TVMFuncRegistry_Get0thFunctionName(reg);

idx = 0;
// NOTE: reg_name_ptr starts at index 1 to skip num_funcs.
for (reg_name_ptr = reg->names + 1; *reg_name_ptr != '\0'; reg_name_ptr++) {
for (; *reg_name_ptr != '\0'; reg_name_ptr++) {
if (!strcmp_cursor(&reg_name_ptr, name)) {
*function_index = idx;
return kTvmErrorNoError;
Expand All @@ -82,9 +97,9 @@ tvm_crt_error_t TVMFuncRegistry_Lookup(const TVMFuncRegistry* reg, const char* n
tvm_crt_error_t TVMFuncRegistry_GetByIndex(const TVMFuncRegistry* reg,
tvm_function_index_t function_index,
TVMBackendPackedCFunc* out_func) {
uint8_t num_funcs;
uint16_t num_funcs;

num_funcs = reg->names[0];
num_funcs = TVMFuncRegistry_GetNumFuncs(reg);
if (function_index >= num_funcs) {
return kTvmErrorFunctionIndexInvalid;
}
Expand All @@ -101,7 +116,8 @@ tvm_crt_error_t TVMMutableFuncRegistry_Create(TVMMutableFuncRegistry* reg, uint8

reg->registry.names = (const char*)buffer;
buffer[0] = 0; // number of functions present in buffer.
buffer[1] = 0; // end of names list marker.
buffer[1] = 0; // note that we combine the first two elements to form a 16-bit function index.
buffer[2] = 0; // end of names list marker.

// compute a guess of the average size of one entry:
// - assume average function name is around ~10 bytes
Expand All @@ -117,13 +133,12 @@ tvm_crt_error_t TVMMutableFuncRegistry_Create(TVMMutableFuncRegistry* reg, uint8
tvm_crt_error_t TVMMutableFuncRegistry_Set(TVMMutableFuncRegistry* reg, const char* name,
TVMBackendPackedCFunc func, int override) {
size_t idx;
char* reg_name_ptr;
char* reg_name_ptr = (char*)TVMFuncRegistry_Get0thFunctionName(&(reg->registry));

idx = 0;
// NOTE: safe to discard const qualifier here, since reg->registry.names was set from
// TVMMutableFuncRegistry_Create above.
// NOTE: reg_name_ptr starts at index 1 to skip num_funcs.
for (reg_name_ptr = (char*)reg->registry.names + 1; *reg_name_ptr != 0; reg_name_ptr++) {
for (; *reg_name_ptr != 0; reg_name_ptr++) {
if (!strcmp_cursor((const char**)&reg_name_ptr, name)) {
if (override == 0) {
return kTvmErrorFunctionAlreadyDefined;
Expand All @@ -149,7 +164,11 @@ tvm_crt_error_t TVMMutableFuncRegistry_Set(TVMMutableFuncRegistry* reg, const ch
reg_name_ptr += name_len + 1;
*reg_name_ptr = 0;
((TVMBackendPackedCFunc*)reg->registry.funcs)[idx] = func;
((char*)reg->registry.names)[0]++; // increment num_funcs.

uint16_t num_funcs;
// increment num_funcs.
num_funcs = TVMFuncRegistry_GetNumFuncs(&(reg->registry)) + 1;
TVMFuncRegistry_SetNumFuncs(&(reg->registry), num_funcs);

return kTvmErrorNoError;
}
Original file line number Diff line number Diff line change
Expand Up @@ -229,7 +229,7 @@ static const TVMBackendPackedCFunc graph_executor_registry_funcs[] = {
};

static const TVMFuncRegistry graph_executor_registry = {
"\x08get_input\0"
"\x08\0get_input\0"
"get_input_index\0"
"get_input_info\0"
"get_num_inputs\0"
Expand Down
8 changes: 7 additions & 1 deletion src/target/func_registry_generator.cc
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,13 @@ namespace target {

std::string GenerateFuncRegistryNames(const Array<String>& function_names) {
std::stringstream ss;
ss << (unsigned char)(function_names.size());

unsigned char function_nums[sizeof(uint16_t)];
*reinterpret_cast<uint16_t*>(function_nums) = function_names.size();
for (auto f : function_nums) {
ss << f;
}

for (auto f : function_names) {
ss << f << '\0';
}
Expand Down
7 changes: 4 additions & 3 deletions tests/crt/func_registry_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ TEST(StrCmpScan, Test) {
}

TEST(FuncRegistry, Empty) {
TVMFuncRegistry registry{"\000", NULL};
TVMFuncRegistry registry{"\000\000", NULL};

EXPECT_EQ(kTvmErrorFunctionNameNotFound, TVMFuncRegistry_Lookup(&registry, "foo", NULL));
EXPECT_EQ(kTvmErrorFunctionIndexInvalid,
Expand All @@ -101,7 +101,7 @@ static int Bar(TVMValue* args, int* type_codes, int num_args, TVMValue* out_ret_
}

// Matches the style of registry defined in generated C modules.
const char* kBasicFuncNames = "\002Foo\0Bar\0"; // NOTE: final \0
const char* kBasicFuncNames = "\002\000Foo\0Bar\0"; // NOTE: final \0
const TVMBackendPackedCFunc funcs[2] = {&Foo, &Bar};
const TVMFuncRegistry kConstRegistry = {kBasicFuncNames, (const TVMBackendPackedCFunc*)funcs};

Expand All @@ -111,7 +111,8 @@ TEST(FuncRegistry, ConstGlobalRegistry) {

// Foo
EXPECT_EQ(kBasicFuncNames[0], 2);
EXPECT_EQ(kBasicFuncNames[1], 'F');
EXPECT_EQ(kBasicFuncNames[1], 0);
EXPECT_EQ(kBasicFuncNames[2], 'F');
EXPECT_EQ(kTvmErrorNoError, TVMFuncRegistry_Lookup(&kConstRegistry, "Foo", &func_index));
EXPECT_EQ(0, func_index);

Expand Down