Python装饰器:将函数转为适配Scikit-learn的DataClass并支持智能提示
解决方案:将Scanpy预处理函数转为Sklearn兼容的Transformer装饰器
核心实现代码
import inspect import dataclasses from sklearn.base import BaseEstimator, TransformerMixin def as_sklearn_transformer(func): # 提取原函数签名与参数信息 sig = inspect.signature(func) param_list = list(sig.parameters.values()) # 过滤无效参数并生成dataclass字段配置 dataclass_fields = {} for param in param_list: # 跳过self/cls这类类方法参数 if param.name in ("self", "cls"): continue field_kwargs = {} # 保留原参数的默认值 if param.default is not inspect.Parameter.empty: field_kwargs["default"] = param.default # 保留原参数的类型注解 if param.annotation is not inspect.Parameter.empty: field_kwargs["type"] = param.annotation dataclass_fields[param.name] = dataclasses.field(**field_kwargs) # 动态生成继承Sklearn基类的dataclass @dataclasses.dataclass class SklearnTransformer(BaseEstimator, TransformerMixin): # 注入生成的dataclass字段 locals().update(dataclass_fields) def fit(self, X, y=None): # 无状态预处理,fit仅返回自身以符合Sklearn规范 return self def transform(self, X): # 提取实例参数,调用原函数 call_args = {k: getattr(self, k) for k in dataclass_fields} # 适配Scanpy的inplace逻辑:默认返回副本而非修改原数据 if "inplace" in sig.parameters: adata_copy = X.copy() func(adata_copy, **call_args) return adata_copy else: return func(X, **call_args) # 复制原函数的文档字符串 SklearnTransformer.__doc__ = func.__doc__ # 设置有意义的类名 SklearnTransformer.__name__ = f"{func.__name__}Transformer" return SklearnTransformer
使用示例
import scanpy as sc from sklearn.pipeline import Pipeline # 将Scanpy的filter_cells转为Sklearn Transformer FilterCells = as_sklearn_transformer(sc.pp.filter_cells) # 实例化时IDE会自动提示具体参数(如min_counts、max_counts) cell_filter = FilterCells(min_counts=100, max_counts=5000) # 集成到Sklearn Pipeline中 preprocessing_pipe = Pipeline([ ("filter_cells", cell_filter), ("normalize", as_sklearn_transformer(sc.pp.normalize_total)) ]) # 测试运行 adata = sc.datasets.pbmc3k() processed_adata = preprocessing_pipe.fit_transform(adata)
问题解决说明
- 智能提示恢复:通过
inspect模块提取原函数的参数名、类型注解与默认值,动态生成dataclass字段,IDE可直接识别具体参数,不再显示**kwargs: Any。 - 字段访问异常修复:利用dataclass自动生成的
__init__方法初始化所有字段,通过getattr访问实例参数时无异常,同时兼容Sklearn的get_params()/set_params()方法。 - Sklearn格式保留:严格继承
BaseEstimator与TransformerMixin,实现标准的fit/transform接口,完全适配Sklearn Pipeline与网格搜索等生态工具。 - Scanpy特性适配:自动处理
inplace参数,默认返回数据副本,符合Sklearn无状态预处理的设计原则,避免修改原始输入数据。
内容的提问来源于stack exchange,提问作者SumNeuron
相关产品推荐
相关产品推荐

