SKLearn自定义OneHotEncoder已调用fit后transform仍提示未拟合如何解决?
问题根因
报错的核心问题出在transform方法的实现上:
- 你写的
OneHotEncoder().transform(X)是每次调用transform时都重新创建了一个全新的、未经过拟合的OneHotEncoder实例,这个新实例自然会抛出未拟合错误。 - 你之前在fit方法中完成拟合的是当前CustomOneHotEncoder的实例本身(也就是
self对象,它继承自OneHotEncoder,已经存储了所有拟合得到的类别信息),但你在transform时完全没有用到这个已拟合的实例。
修复方案
直接修改transform方法的调用逻辑,调用当前已拟合实例的父类transform方法即可,修改后的代码如下:
import numpy as np import pandas as pd from sklearn.preprocessing import OneHotEncoder class CustomOneHotEncoder(OneHotEncoder): """ OneHot Encoding ---------- """ # todo: max_num_categories 作为参数 def __init__(self, categories='auto', drop=None, sparse_output=True, dtype=np.int32, handle_unknown='error'): # 注意:sklearn 1.2及以上版本稀疏输出参数名从sparse改为sparse_output,如果你用的是旧版本可以改回sparse super().__init__( categories=categories, drop=drop, sparse_output=sparse_output, dtype=dtype, handle_unknown=handle_unknown ) # 如果你没有额外的fit逻辑,这个fit方法可以直接省略,父类的实现已经满足要求 def fit(self, X, y=None): super().fit(X, y=y) return self def transform(self, X): """ :type X: DataFrame """ try: # 调用父类的transform方法,使用当前已拟合的实例参数 ret = super().transform(X).toarray() return ret except Exception as e: # 建议保留原始错误信息方便调试,不要直接吞掉原始异常 raise Exception(f"Internal Error: {str(e)}") from e
额外说明
- 如果你的自定义转换器没有额外的业务逻辑,其实不需要重写fit和__init__方法,直接复用父类的实现即可,只需要保留你要自定义的transform逻辑就行。
- 原代码中异常捕获会吞掉原始报错信息,不利于问题排查,建议修改为带上原始异常的形式。
内容的提问来源于stack exchange,提问作者gilgamash
相关产品推荐
相关产品推荐

