元编程实践:使用__class__时如何从模板生成自定义类?
嘿,这个问题我刚好有实用的解决思路!你完全不用重复写那些大同小异的类,用Python的类工厂函数这种轻量元编程技巧就能搞定——比Jinja模板生成代码更灵活,还能在运行时直接生成可用的类,不需要额外的代码生成步骤。
核心思路:用工厂函数动态生成类
我们可以写一个通用的工厂函数,接收你想要扩展的sklearn父类(比如VarianceThreshold、SelectKBest等),然后动态生成带有feature_names追踪功能的子类。原来代码里需要修改的两处硬编码,都会被替换成动态传入的父类参数。
完整实现代码
import copy def create_feature_name_tracker(base_class): """动态生成带有feature_names追踪功能的sklearn子类""" class FeatureNameTracker(base_class): def __init__(self, **kwargs): super().__init__(**kwargs) self.feature_names = None def get_params(self, deep=True): # 保留你原来的hack逻辑,只是把硬编码的类换成传入的base_class params = super().get_params(deep) cp = copy.copy(self) cp.__class__ = base_class params.update(cp.__class__.get_params(cp, deep)) return params def fit(self, X, y=None): self.feature_names = list(X.columns) return super().fit(X, y) # 给生成的类起个符合你习惯的名字,比如VarianceThresholdN、SelectKBestN FeatureNameTracker.__name__ = f"{base_class.__name__}N" return FeatureNameTracker
怎么用这个工厂函数?
只需要传入你想要扩展的sklearn类,就能得到对应的自定义类:
from sklearn.feature_selection import VarianceThreshold, SelectKBest, f_regression import pandas as pd # 生成你原来的VarianceThresholdN VarianceThresholdN = create_feature_name_tracker(VarianceThreshold) # 生成另一个自定义类SelectKBestN SelectKBestN = create_feature_name_tracker(SelectKBest) # 测试功能 X = pd.DataFrame({'a': [1,2,3], 'b': [4,5,6], 'c': [7,8,9]}) y = [0,1,0] # 测试VarianceThresholdN vt = VarianceThresholdN(threshold=0) vt.fit(X) print(vt.feature_names) # 输出: ['a', 'b', 'c'] print(vt.get_params()) # 能正确获取父类的所有参数(比如threshold) # 测试SelectKBestN skb = SelectKBestN(score_func=f_regression, k=2) skb.fit(X, y) print(skb.feature_names) # 输出: ['a', 'b', 'c'] print(skb.get_params()) # 能正确获取score_func、k等参数
为什么这个方法比重复写类好?
- 完全避免重复代码:不管你要扩展多少个sklearn类,只需要调用一次工厂函数就行,不用复制粘贴修改两处硬编码。
- 兼容sklearn生态:保留了你原来的
get_params逻辑,生成的类能正常用于sklearn的管道、网格搜索等组件,不会出现参数识别问题。 - 灵活可控:如果以后要修改
feature_names的逻辑,只需要改工厂函数里的代码,所有生成的类都会自动更新。
内容的提问来源于stack exchange,提问作者Quant Christo
相关产品推荐
相关产品推荐

