如何在工厂模式中正确实现带专属函数的派生类?
工厂模式下派生类专属函数的面向对象设计方案
针对工厂模式返回基类指针后无法便捷调用派生类专属函数、基类虚函数臃肿扩展性差的问题,以下是几种贴合C++场景的解决方案:
方案一:安全类型转换(dynamic_cast)
直接通过dynamic_cast将基类指针转换为目标派生类指针,转换后检查指针有效性,确保调用专属函数的安全性。这种方式适合明确知道对象实际类型的场景,无需修改原有类结构。
代码示例
#include <iostream> #include <string> #include <vector> // 基类 class Inference { public: virtual ~Inference() = default; // 虚析构确保正确销毁 }; // 派生类Detector,带专属初始化函数 class Detector : public Inference { public: void Initialize(std::string model_path, std::vector<std::string> output_info) { std::cout << "Detector初始化:模型路径" << model_path << ",输出信息:"; for (const auto& info : output_info) { std::cout << info << " "; } std::cout << std::endl; } }; // 派生类Classifier,带专属初始化函数 class Classifier : public Inference { public: void Initialize(std::string model_path) { std::cout << "Classifier初始化:模型路径" << model_path << std::endl; } }; // 工厂类 class InferenceFactory { public: static Inference* Create(const std::string& type) { if (type == "Detector") { return new Detector(); } else if (type == "Classifier") { return new Classifier(); } return nullptr; } }; // 使用示例 int main() { Inference* infer = InferenceFactory::Create("Detector"); if (Detector* detector = dynamic_cast<Detector*>(infer)) { detector->Initialize("/models/yolov8.pt", {"bounding_box", "label", "conf"}); } else { std::cerr << "类型转换失败,无法调用Detector专属函数" << std::endl; } Inference* infer2 = InferenceFactory::Create("Classifier"); if (Classifier* classifier = dynamic_cast<Classifier*>(infer2)) { classifier->Initialize("/models/resnet50.pt"); } else { std::cerr << "类型转换失败,无法调用Classifier专属函数" << std::endl; } delete infer; delete infer2; return 0; }
注意:使用dynamic_cast需要基类包含至少一个虚函数(如虚析构),且运行时类型信息(RTTI)需开启(多数编译器默认开启)。
方案二:接口分离原则(ISP)
将不同派生类的专属功能拆分为独立的抽象接口,派生类同时继承基类和对应的专属接口。工厂返回基类指针后,按需转换到专属接口调用函数,避免基类被无关虚函数污染。
代码示例
#include <iostream> #include <string> #include <vector> // 基础推理接口(仅包含所有派生类通用方法) class Inference { public: virtual ~Inference() = default; virtual void RunInference(const std::vector<float>& input) = 0; // 通用推理方法 }; // Detector专属初始化接口 class IDetectorInit { public: virtual ~IDetectorInit() = default; virtual void Initialize(std::string model_path, std::vector<std::string> output_info) = 0; }; // Classifier专属初始化接口 class IClassifierInit { public: virtual ~IClassifierInit() = default; virtual void Initialize(std::string model_path) = 0; }; // Detector实现两个接口 class Detector : public Inference, public IDetectorInit { public: void Initialize(std::string model_path, std::vector<std::string> output_info) override { std::cout << "Detector初始化完成:" << model_path << std::endl; } void RunInference(const std::vector<float>& input) override { std::cout << "Detector执行推理,输入维度:" << input.size() << std::endl; } }; // Classifier实现两个接口 class Classifier : public Inference, public IClassifierInit { public: void Initialize(std::string model_path) override { std::cout << "Classifier初始化完成:" << model_path << std::endl; } void RunInference(const std::vector<float>& input) override { std::cout << "Classifier执行推理,输入维度:" << input.size() << std::endl; } }; // 工厂类 class InferenceFactory { public: static Inference* Create(const std::string& type) { if (type == "Detector") { return new Detector(); } else if (type == "Classifier") { return new Classifier(); } return nullptr; } }; // 使用示例 int main() { Inference* infer = InferenceFactory::Create("Detector"); if (IDetectorInit* detector_init = dynamic_cast<IDetectorInit*>(infer)) { detector_init->Initialize("/models/yolov8.pt", {"bounding_box", "label"}); } infer->RunInference({1.0f, 2.0f, 3.0f}); Inference* infer2 = InferenceFactory::Create("Classifier"); if (IClassifierInit* classifier_init = dynamic_cast<IClassifierInit*>(infer2)) { classifier_init->Initialize("/models/resnet50.pt"); } infer2->RunInference({4.0f, 5.0f, 6.0f}); delete infer; delete infer2; return 0; }
这种方式遵循单一职责原则,每个接口只负责一类功能,新增派生类时只需定义对应的专属接口,无需修改现有类。
方案三:统一配置类初始化
针对初始化参数不同的场景,定义包含所有可能参数的配置类,基类提供统一的初始化接口,派生类根据自身需求从配置类中提取参数。
代码示例
#include <iostream> #include <string> #include <vector> #include <optional> // 统一配置类 struct InferenceConfig { std::string model_path; std::optional<std::vector<std::string>> detector_output_info; // Detector专属可选参数 // 可扩展其他派生类的专属参数 }; // 基类 class Inference { public: virtual ~Inference() = default; virtual void Initialize(const InferenceConfig& config) = 0; virtual void RunInference(const std::vector<float>& input) = 0; }; class Detector : public Inference { public: void Initialize(const InferenceConfig& config) override { if (!config.detector_output_info.has_value()) { std::cerr << "Detector初始化缺少输出信息参数" << std::endl; return; } std::cout << "Detector初始化:模型路径" << config.model_path << ",输出信息:"; for (const auto& info : config.detector_output_info.value()) { std::cout << info << " "; } std::cout << std::endl; } void RunInference(const std::vector<float>& input) override { std::cout << "Detector执行推理" << std::endl; } }; class Classifier : public Inference { public: void Initialize(const InferenceConfig& config) override { std::cout << "Classifier初始化:模型路径" << config.model_path << std::endl; // Classifier无需detector_output_info,直接忽略 } void RunInference(const std::vector<float>& input) override { std::cout << "Classifier执行推理" << std::endl; } }; // 工厂类 class InferenceFactory { public: static Inference* Create(const std::string& type) { if (type == "Detector") { return new Detector(); } else if (type == "Classifier") { return new Classifier(); } return nullptr; } }; // 使用示例 int main() { InferenceConfig detector_config{ .model_path = "/models/yolov8.pt", .detector_output_info = {"bounding_box", "label", "conf"} }; Inference* detector = InferenceFactory::Create("Detector"); detector->Initialize(detector_config); InferenceConfig classifier_config{ .model_path = "/models/resnet50.pt" }; Inference* classifier = InferenceFactory::Create("Classifier"); classifier->Initialize(classifier_config); delete detector; delete classifier; return 0; }
这种方式避免了重载多个初始化函数,新增派生类时只需在配置类中添加对应可选参数,原有代码无需修改,扩展性强。
方案四:访问者模式
如果需要对不同派生类执行多种差异化操作,可使用访问者模式。基类定义Accept方法接收访问者,访问者类针对每个派生类实现对应的操作逻辑,无需在基类中添加新的虚函数。
代码示例
#include <iostream> #include <string> #include <vector> #include <optional> // 前置声明访问者类 class InferenceVisitor; // 基类 class Inference { public: virtual ~Inference() = default; virtual void Accept(InferenceVisitor& visitor) = 0; }; // 派生类Detector class Detector : public Inference { public: std::string model_path; std::vector<std::string> output_info; void Accept(InferenceVisitor& visitor) override; // Detector专属方法 void SetOutputInfo(const std::vector<std::string>& info) { output_info = info; } }; // 派生类Classifier class Classifier : public Inference { public: std::string model_path; void Accept(InferenceVisitor& visitor) override; // Classifier专属方法 void SetModelPrecision(float precision) { std::cout << "Classifier设置精度:" << precision << std::endl; } }; // 访问者类 class InferenceVisitor { public: virtual void Visit(Detector& detector) = 0; virtual void Visit(Classifier& classifier) = 0; }; // 初始化访问者 class InitVisitor : public InferenceVisitor { private: std::string model_path_; std::optional<std::vector<std::string>> detector_output_; float classifier_precision_; public: InitVisitor(std::string model_path) : model_path_(model_path) {} InitVisitor(std::string model_path, std::vector<std::string> output) : model_path_(model_path), detector_output_(output) {} InitVisitor(std::string model_path, float precision) : model_path_(model_path), classifier_precision_(precision) {} void Visit(Detector& detector) override { detector.model_path = model_path_; if (detector_output_.has_value()) { detector.SetOutputInfo(detector_output_.value()); } std::cout << "Detector初始化完成" << std::endl; } void Visit(Classifier& classifier) override { classifier.model_path = model_path_; classifier.SetModelPrecision(classifier_precision_); std::cout << "Classifier初始化完成" << std::endl; } }; // 实现Accept方法 void Detector::Accept(InferenceVisitor& visitor) { visitor.Visit(*this); } void Classifier::Accept(InferenceVisitor& visitor) { visitor.Visit(*this); } // 工厂类 class InferenceFactory { public: static Inference* Create(const std::string& type) { if (type == "Detector") { return new Detector(); } else if (type == "Classifier") { return new Classifier(); } return nullptr; } }; // 使用示例 int main() { Inference* detector = InferenceFactory::Create("Detector"); InitVisitor detector_init("/models/yolov8.pt", {"bounding_box", "label"}); detector->Accept(detector_init); Inference* classifier = InferenceFactory::Create("Classifier"); InitVisitor classifier_init("/models/resnet50.pt", 0.9f); classifier->Accept(classifier_init); delete detector; delete classifier; return 0; }
访问者模式适合操作种类较多且易扩展的场景,新增操作只需添加新的访问者类,无需修改原有派生类。
内容的提问来源于stack exchange,提问作者dungdq
相关产品推荐
相关产品推荐

