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);
|
问题:
- 维护成本高:添加新算子需要修改两个地方。
- 一致性问题:Operator 和 Kernel 的注册可能不同步。
- 查找开销:执行时需要先查 Operator,再查 Kernel。
新架构:统一的 PluginRegistry,一次注册,直接使用。
2. PluginKey 设计
A. 二元组 Key
1 2 3 4 5 6 7 8 9 10
|
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)); 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>()); } };
|
工作原理:
PluginRegistrar 是一个模板类。
- 构造函数中调用
register_creator。
- 创建全局静态实例时,构造函数自动执行。
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
| 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
| auto plugin = PluginRegistry::instance().create_plugin( node->type(), config_.device_type);
auto* generic_op = dynamic_cast<GenericOperator*>(node->get_operator().get()); if (generic_op && generic_op->plugin_param()) { plugin->set_param(generic_op->plugin_param()); }
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
| 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 推理的端到端示例。