如何遵循SOLID原则动态调用数据集加载类?
问题背景
我正在开发一个遵循SOLID原则的数据集生成模块,以下是基于**里氏替换原则(Liskov Substitution Principle)**设计的代码:
from abc import ABC class BaseLoader(ABC): def __init__(self, dataset_name='mnist'): self.dataset_name = dataset_name class MNISTLoader(BaseLoader): def load(self): # 加载数据的逻辑 pass class OCTMNISTLoader(BaseLoader): def download(self): # 下载数据的逻辑 pass
现在我希望根据解析后的参数或加载的配置文件动态创建类实例,想知道以下实现是否为最佳实践,或者有没有更优的动态实例化方式:
possible_instances = {'mnist': MNISTLoader, 'octmnist': OCTMNISTLoader} chosen_dataset = 'mnist' instance = possible_instances[chosen_dataset](dataset_name=chosen_dataset)
补充说明
我们还考虑过用函数来动态调用类,把这个函数放在包含这些类的模块中:
def get_loader(loader_name: str) -> BaseLoader: loaders = { 'mnist': MNISTLoader, 'octmnist': OCTMNISTLoader } try: return loaders[loader_name] except KeyError as err: raise CustomError("合适的错误提示信息")
我不确定哪种方式最符合Python风格。
分析与建议
两种方式的优劣对比
直接字典映射实例化
- 优点:代码简洁直观,适合小规模、场景简单的情况,无需额外封装函数,直接通过字典查找类并实例化,可读性强。
- 缺点:后续如果需要扩展(比如添加参数校验、统一初始化逻辑、错误处理),相关代码会分散在业务逻辑中,导致冗余;
KeyError的捕获需要在调用处单独处理,无法统一管理。
封装成工厂函数
- 优点:符合单一职责原则,将加载器的查找、实例化、错误处理逻辑集中在一处,业务代码只需调用
get_loader即可,降低耦合;可以统一抛出自定义异常、提供友好提示,后续新增加载器只需修改工厂函数内的字典,无需改动所有调用位置;还能在函数内添加额外逻辑,比如参数校验、统一初始化参数处理。 - 缺点:对于极简单的场景会多一层封装,但长期来看更利于维护。
- 优点:符合单一职责原则,将加载器的查找、实例化、错误处理逻辑集中在一处,业务代码只需调用
更符合Python风格的选择
如果你的模块后续有扩展需求(比如新增更多数据集加载器)、需要统一的错误处理或初始化逻辑,工厂函数的方式更符合Python“显式优于隐式”“DRY(Don't Repeat Yourself)”的设计原则,是更推荐的实践。
另外,建议修正工厂函数的细节:如果目标是返回类实例而非类本身,应在返回时添加实例化括号并传递必要参数,示例如下:
def get_loader(loader_name: str, dataset_name: str) -> BaseLoader: loaders = { 'mnist': MNISTLoader, 'octmnist': OCTMNISTLoader } try: return loaders[loader_name](dataset_name=dataset_name) except KeyError: raise CustomError(f"不支持的数据集加载器:{loader_name}")
进阶优化建议
如果后续加载器数量较多,可以用装饰器自动注册加载器,避免每次新增加载器都手动修改工厂函数内的字典:
from abc import ABC loaders_registry = {} def register_loader(name: str): def decorator(cls): loaders_registry[name] = cls return cls return decorator class BaseLoader(ABC): def __init__(self, dataset_name='mnist'): self.dataset_name = dataset_name @register_loader('mnist') class MNISTLoader(BaseLoader): def load(self): pass @register_loader('octmnist') class OCTMNISTLoader(BaseLoader): def download(self): pass def get_loader(loader_name: str, dataset_name: str) -> BaseLoader: try: return loaders_registry[loader_name](dataset_name=dataset_name) except KeyError: raise CustomError(f"不支持的数据集加载器:{loader_name}")
这种方式灵活性更强,新增加载器只需添加装饰器,无需修改工厂函数,更贴合Python的动态特性。
内容的提问来源于stack exchange,提问作者MaKaNu
相关产品推荐
相关产品推荐

