You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)

问题解决说明

  1. 智能提示恢复:通过inspect模块提取原函数的参数名、类型注解与默认值,动态生成dataclass字段,IDE可直接识别具体参数,不再显示**kwargs: Any。
  2. 字段访问异常修复:利用dataclass自动生成的__init__方法初始化所有字段,通过getattr访问实例参数时无异常,同时兼容Sklearn的get_params()/set_params()方法。
  3. Sklearn格式保留:严格继承BaseEstimator与TransformerMixin,实现标准的fit/transform接口,完全适配Sklearn Pipeline与网格搜索等生态工具。
  4. Scanpy特性适配:自动处理inplace参数,默认返回数据副本,符合Sklearn无状态预处理的设计原则,避免修改原始输入数据。

内容的提问来源于stack exchange,提问作者SumNeuron

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 02:17:13