使用它包装的函数保存sklearn`PortTransformer`

Uri*_*ren 7 python pickle scikit-learn joblib

我使用sklearn的Pipeline,并FunctionTransformer用自定义功能

from sklearn.externals import joblib
from sklearn.preprocessing import FunctionTransformer
from sklearn.pipeline import Pipeline
Run Code Online (Sandbox Code Playgroud)

这是我的代码:

def f(x):
    return x*2
pipe = Pipeline([("times_2", FunctionTransformer(f))])
joblib.dump(pipe, "pipe.joblib")
del pipe
del f
pipe = joblib.load("pipe.joblib") # Causes an exception
Run Code Online (Sandbox Code Playgroud)

我收到这个错误:

AttributeError:模块'__ main__'没有属性'f'

怎么解决这个问题?

请注意,此问题也发生在 pickle

Uri*_*ren 5

我能够使用该marshal模块破解解决方案(除此之外pickle)并覆盖魔术方法getstate并setstate使用pickle.

import marshal
from types import FunctionType
from sklearn.base import BaseEstimator, TransformerMixin

class MyFunctionTransformer(BaseEstimator, TransformerMixin):
    def __init__(self, f):
        self.func = f
    def __call__(self, X):
        return self.func(X)
    def __getstate__(self):
        self.func_name = self.func.__name__
        self.func_code = marshal.dumps(self.func.__code__)
        del self.func
        return self.__dict__
    def __setstate__(self, d):
        d["func"] = FunctionType(marshal.loads(d["func_code"]), globals(), d["func_name"])
        del d["func_name"]
        del d["func_code"]
        self.__dict__ = d
    def fit(self, X, y=None):
        return self
    def transform(self, X):
        return self.func(X)
Run Code Online (Sandbox Code Playgroud)

现在,如果我们使用MyFunctionTransformer而不是FunctionTransformer代码,代码按预期工作:

from sklearn.externals import joblib
from sklearn.pipeline import Pipeline

@MyFunctionTransformer
def my_transform(x):
    return x*2
pipe = Pipeline([("times_2", my_transform)])
joblib.dump(pipe, "pipe.joblib")
del pipe
del my_transform
pipe = joblib.load("pipe.joblib")
Run Code Online (Sandbox Code Playgroud)

这种方式的工作方式是删除fpickle中的函数,而不是marshaling它的代码和名称.

dill 对于编组而言,它看起来也是一个很好的选择