Mini-Infer (34): 插件架构 (下) — PluginRegistry 与自动注册

1. 问题背景:双层注册的痛点

在旧架构中,我们有两个独立的注册表:

1
2
3
4
// 旧架构
OperatorFactory::register_operator("Conv2D", Conv2DOperatorCreator);
KernelRegistry::register_kernel("Conv2D", CPU, Conv2DCPUKernel);
KernelRegistry::register_kernel("Conv2D", CUDA, Conv2DCUDAKernel);

问题:

  1. 维护成本高:添加新算子需要修改两个地方。
  2. 一致性问题:Operator 和 Kernel 的注册可能不同步。
  3. 查找开销:执行时需要先查 Operator,再查 Kernel。

新架构:统一的 PluginRegistry,一次注册,直接使用。


2. PluginKey 设计

A. 二元组 Key

1
2
3
4
5
6
7
8
9
10
// mini_infer/operators/plugin_registry.h

struct PluginKey {
core::OpType op_type;
core::DeviceType device_type;

bool operator==(const PluginKey& other) const {
return op_type == other.op_type && device_type == other.device_type;
}
};

Key = (OpType, DeviceType):

  • OpType:算子类型(Conv2D、ReLU、Pooling 等)。
  • DeviceType:设备类型(CPU、CUDA)。

为什么不包含 DataType?

  • 大多数算子支持多种数据类型。
  • 数据类型在运行时由 Plugin 内部处理。
  • 减少注册表的复杂度。

B. PluginKeyHash 实现

1
2
3
4
5
6
7
8
9
10
struct PluginKeyHash {
size_t operator()(const PluginKey& key) const {
const auto op_hash = std::hash<int>{}(static_cast<int>(key.op_type));
const auto dev_hash = std::hash<int>{}(static_cast<int>(key.device_type));
// 使用 boost::hash_combine 风格的哈希组合
size_t seed = op_hash;
seed ^= dev_hash + 0x9e3779b9 + (seed << 6) + (seed >> 2);
return seed;
}
};

0x9e3779b9 是黄金比例的倒数的 32 位表示,用于减少哈希冲突。


3. PluginRegistry 单例模式

A. 类定义

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class PluginRegistry {
public:
static PluginRegistry& instance() {
static PluginRegistry registry;
return registry;
}

void register_creator(std::unique_ptr<IPluginCreator> creator);
std::unique_ptr<IPlugin> create_plugin(core::OpType op_type,
core::DeviceType device_type) const;
bool has_plugin(core::OpType op_type, core::DeviceType devpe) const;
std::vector<PluginKey> get_registered_keys() const;

private:
PluginRegistry() = default;
PluginRegistry(const PluginRegistry&) = delete;
PluginRegistry& operator=(const PluginRegistry&) = delete;

mutable std::mutex mutex_;
std::unordered_map<PluginKey, std::unique_ptr<IPluginCreator>, PluginKeyHash> creators_;
};

B. register_creator 实现

1
2
3
4
5
6
7
8
9
void PluginRegistry::register_creator(std::unique_ptr<IPluginCreator> creator) {
if (!creator) {
return;
}

PluginKey key{creator->get_op_type(), creator->get_device_type()};
std::lock_guard<std::mutex> lock(mutex_);
creators_[key] = std::move(creator);
}

线程安全:使用 std::mutex 保护注册操作。

C. create_plugin 实现

1
2
3
4
5
6
7
8
9
10
11
12
std::unique_ptr<IPlugin> PluginRegistry::create_plugin(
core::OpType op_type,
core::DeviceType device_type) const {

PluginKey key{op_type, device_type};
std::lock_guard<std::mutex> lock(mutex_);
auto it = creators_.find(key);
if (it == creators_.end()) {
return nullptr;
}
return it->second->create_plugin();
}

使用示例:

1
2
auto plugin = PluginRegistry::instance().create_plugin(
core::OpType::kCONV2D, core::DeviceType::CUDA);

D. has_plugin / get_registered_keys

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
bool PluginRegistry::has_plugin(core::OpType op_type,
core::DeviceType device_type) const {
PluginKey key{op_type, device_type};
std::lock_guard<std::mutex> lock(mutex_);
return creators_.find(key) != creators_.end();
}

std::vector<PluginKey> PluginRegistry::get_registered_keys() const {
std::lock_guard<std::mutex> lock(mutex_);
std::vector<PluginKey> keys;
keys.reserve(creators_.size());
for (const auto& pair : creators_) {
keys.push_back(pair.first);
}
return keys;
}

4. 静态注册机制

A. PluginRegistrar 模板

1
2
3
4
5
6
7
template <typename CreatorType>
class PluginRegistrar {
public:
PluginRegistrar() {
PluginRegistry::instance().register_creator(std::make_unique<CreatorType>());
}
};

工作原理:

  1. PluginRegistrar 是一个模板类。
  2. 构造函数中调用 register_creator。
  3. 创建全局静态实例时,构造函数自动执行。

B. REGISTER_PLUGIN 宏

1
2
3
4
5
#define REGISTER_PLUGIN(plugin_class, creator_class) \
namespace { \
static ::mini_infer::operators::PluginRegistrar<creator_class> \
g_##plugin_class##_registrar; \
}

展开示例:

1
2
3
4
5
6
7
REGISTER_PLUGIN(ReLUCPUPlugin, ReLUCPUPluginCreator)

// 展开为:
namespace {
static ::mini_infer::operators::PluginRegistrar<ReLUCPUPluginCreator>
g_ReLUCPUPlugin_registrar;
}

匿名命名空间:确保静态变量只在当前编译单元可见,避免链接冲突。

C. DEFINE_PLUGIN_CREATOR 宏

1
2
3
4
5
6
7
8
9
10
11
12
13
14
#define DEFINE_PLUGIN_CREATOR(plugin_class, type_name, op_type_enum, device_type_enum) \
class plugin_class##Creator : public ::mini_infer::operators::IPluginCreator { \
public: \
const char* get_plugin_type() const noexcept override { return type_name; } \
::mini_infer::core::OpType get_op_type() const noexcept override { \
return ::mini_infer::core::OpType::op_type_enum; \
} \
::mini_infer::core::DeviceType get_device_type() const noexcept override { \
return ::mini_:DeviceType::device_type_enum; \
} \
std::unique_ptr<::mini_infer::operators::IPlugin> create_plugin() const override { \
return std::make_unique<plugin_class>(); \
} \
};

展开示例:

1
2
3
4
5
6
7
8
9
10
11
12
DEFINE_PLUGIN_CREATOR(ReLUCPUPlugin, "Relu", kRELU, CPU)

// 展开为:
class ReLUCPUPluginCreator : public IPluginCreator {
public:
const char* get_plugin_type() const noexcept override {Relu"; }
core::OpType get_op_type() const noexcept override { return core::OpType::kRELU; }
core::DeviceType get_device_type() const noexcept override { return core::DeviceType::CPU; }
std::unique_ptr<IPlugin> create_plugin() const override {
return std::make_unique<ReLUCPUPlugin>();
}
};

D. REGISTER_PLUGIN_SIMPLE 组合宏

1
2
3
#define REGISTER_PLUGIN_SIMPLE(plugin_class, type_name, op_type_enum, device_type_enum) \
DEFINE_PLUGIN_CREATOR(plugin_class, type_name, op_type_enum, device_type_enum) \
REGISTER_PLUGIN(plugin_class, plugin_class##Creator)

一行代码完成定义和注册:

1
2
// 在 relu_cpu.cpp 末尾
REGISTER_PLUGIN_SIMPLE(ReLUCPUPlugin, "Relu", kRELU, CPU)

5. 与 ONNX 导入的集成

A. GenericOperator 存储参数

当从 ONNX 导入模型时,我们使用 GenericOperator 存储算子参数:

1
2
3
4
5
6
7
8
9
10
11
12
13
class GenericOperator : public Operator {
public:
void set_plugin_param(std::shared_ptr<PluginParam> param) {
plugin_param_ = param;
}

std::shared_ptr<PluginParam> plugin_param() const {
return plugin_param_;
}

private:
std::shared_ptr<PluginParam> plugin_param_;
};

B. InferencePlan 创建并缓存 Plugin

在 InferencePlan::infer_shapes 中:

1
2
3
4
5
6
7
8
9
10
11
12
// 创建 Plugin
auto plugin = PluginRegistry::instance().create_plugin(
node->type(), config_.device_type);

// 从 GenericOperator 传递参数
auto* generic_op = dynamic_cast<GenericOperator*>(node->get_operator().get());
if (generic_op && generic_op->plugin_param()) {
plugin->set_param(generic_op->plugin_param());
}

// 缓存 Plugin 供执行时使用
node->get_operator()->set_cached_plugin(std::move(plugin));

6. 链接器的"副作用"驱动

A. 静态初始化顺序

C++ 的静态变量在 main() 之前初始化。我们利用这一特性实现自动注册:

1
2
3
4
5
6
7
8
9
10
11
程序启动
↓
静态变量初始化
↓
PluginRegistrar 构造函数执行
↓
register_creator 被调用
↓
Plugin 注册到 Registry
↓
main() 开始执行

B. 确保注册代码被链接

问题:如果 Plugin 实现在静态库中,且没有被显式引用,链接器可能不会包含它。

解决方案 1:使用 --whole-archive(Linux)或 /WHOLEARCHIVE(Windows)。

1
2
3
4
5
target_link_libraries(my_app
-Wl,--whole-archive
mini_infer_plugins
-Wl,--no-whole-archive
)

解决方案 2:在主程序中显式引用。

1
2
3
4
5
6
7
8
9
// 在 main.cpp 中
extern void force_link_cpu_plugins();
extern void force_link_cuda_plugins();

int main() {
force_link_cpu_plugins();
force_link_cuda_plugins();
// ...
}

解决方案 3:使用 CMake 的 OBJECT 库。

1
2
3
4
5
6
7
add_library(plugins OBJECT
relu_cpu.cpp
conv2d_cpu.cpp
# ...
)

target_link_libraries(my_app PRIVATE $<TARGET_OBJECTS:plugins>)

7. 完整注册流程图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
┌─────────────────────────────────────────────────────────────────┐
│ 编译时 │
├─────────────────────────────────────────────────────────────────┤
│ relu_cpu.cpp: │
│ class ReLUCPUPlugin : public SimpleCPUPlugin<...> { ... } │
│ REGISTER_PLUGIN_SIMPLE(ReLUCPUPlugin, "Relu", kRELU, CPU) │
│ ↓ │
│ 展开为: │
│ - class ReLUCPUPluginCreator : public IPluginCreator { ... } │
│ - static PluginRegistrar<ReLUCPUPluginCreator> g_...; │
└─────────────────────────────────────────────────────────────────┘
↓
┌─────────────────────────────────────────────────────────────────┐
│ 运行时(main 之前) │
├─────────────────────────────────────────────────────────────────┤
│ 静态变量 g_ReLUCPUPlugin_registrar 初始化 │
│ ↓ │
│ PluginRegistrar 构造函数执行 │
│ ↓ │
│ PluginRegistry::instance().register_creator( │
│ std::make_unique<ReLUCPUPluginCreator>()) │
│ ↓ │
│ creators_[{kRELU, CPU}] = ReLUCPUPluginCreator │
└─────────────────────────────────────────────────────────────────┘
↓
┌─────────────────────────────────────────────────────────────────┐
│ 运行时(main 之后) │
├─────────────────────────────────────────────────────────────────┤
│ auto plugin = PluginRegistry::instance().create_plugin( │
│ kRELU, CPU); │
│ ↓ │
│ creators_[{kRELU, CPU}]->create_plugin() │
│ ↓ │
│ return std::make_unique<ReLUCPUPlugin>() │
└─────────────────────────────────────────────────────────────────┘

8. 总结

本篇我们实现了 Mini-Infer 插件架构的注册机制:

  • PluginKey:(OpType, DeviceType) 二元组作为注册表 Key。
  • PluginRegistry 单例:线程安全的全局注册表。
  • 静态注册机制:利用 C++ 静态初始化实现自动注册。
  • 宏简化:REGISTER_PLUGIN_SIMPLE 一行代码完成注册。
  • 链接器注意事项:确保静态库中的注册代码被链接。

下一篇,我们将完成插件架构的最后一部分——从旧架构到新架构的迁移实战,以及 CUDA 推理的端到端示例。