如何获取Spark ML OneHotEncoder生成的编码对应类别名称
问题背景
我在Spark ML流水线中使用了OneHotEncoder,相关实现代码如下:
from pyspark.ml.feature import StringIndexer, OneHotEncoder, VectorAssembler schema = StructType( [StructField("PassengerId", DoubleType()), StructField("Survived", DoubleType()), StructField("Pclass", DoubleType()), StructField("Name", StringType()), StructField("Sex", StringType()), StructField("Age", DoubleType()), StructField("SibSp", DoubleType()), StructField("Parch", DoubleType()), StructField("Ticket", StringType()), StructField("Fare", DoubleType()), StructField("Cabin", StringType()), StructField("Embarked", StringType()) ]) titanic = ( spark .read .option("header", "true") .schema(schema) .csv("/working_dir/data/titanic.csv") .na.fill(0) .na.fill("Nulo") ) trainDF, testDF = titanic.randomSplit([0.8, 0.2], seed=42) indexers = [StringIndexer(inputCol=c, outputCol=c + "_index", handleInvalid="keep") for c in ["Sex", "Embarked"]] ohe = OneHotEncoder(inputCols=[indexer.getOutputCol() for indexer in indexers], outputCols=[f"{indexer.getInputCol()}_onehot" for indexer in indexers]) vectorAssembler = ( VectorAssembler() .setInputCols(["Sex_onehot", "Embarked_onehot", "Pclass", "Age", "Fare"]) .setOutputCol("features") ) from pyspark.ml import Pipeline pipeline = Pipeline(stages=[ *indexers, ohe, vectorAssembler ]) trainDF, testDF = titanic.randomSplit([0.8, 0.2], seed=42) fitted_pipeline = pipeline.fit(trainDF)
目前需要提取每个变量经编码后生成的类别名称,希望实现和sklearn中如下调用类似的输出效果:
>>> enc.get_feature_names_out(['gender', 'group']) array(['gender_Female', 'gender_Male', 'group_1', 'group_2', 'group_3'], ...)
此前尝试打印ParamMap提取相关类别信息,未取得预期结果,询问可行实现方法。
实现方案
Spark ML的OneHotEncoder默认采用丢弃最后一个类别的独热编码逻辑(避免多重共线性),编码后的类别映射信息保存在拟合后的StringIndexer模型中,按以下步骤即可拿到和sklearn一致格式的特征名:
- 从拟合完成的pipeline中提取各个拟合好的StringIndexer模型,拿到每个原始分类列对应的类别顺序
- 匹配OneHotEncoder的输入输出列映射,按独热编码的生成规则拼接特征名
- 最后拼接数值类特征的列名,就能得到完整的features列对应的所有特征名称
可直接复用的实现代码如下:
from pyspark.ml.feature import StringIndexerModel, OneHotEncoderModel # 从拟合后的流水线中提取各阶段模型 fitted_stages = fitted_pipeline.stages # 提取所有StringIndexerModel,存为 原始列名: 类别列表 的映射 indexer_category_map = {} for stage in fitted_stages: if isinstance(stage, StringIndexerModel): original_col = stage.getInputCol() # labels属性就是按索引排序的类别值,索引0对应第一个类别,以此类推 indexer_category_map[original_col] = stage.labels # 提取拟合后的OneHotEncoderModel ohe_model = [s for s in fitted_stages if isinstance(s, OneHotEncoderModel)][0] ohe_feature_names = [] # 遍历每个独热编码的输入输出列对 for in_col, out_col in zip(ohe_model.getInputCols(), ohe_model.getOutputCols()): # 从索引列名反推原始分类列名:去掉_index后缀 original_col = in_col.replace("_index", "") categories = indexer_category_map[original_col] # dropLast模式下会丢弃最后一个类别,所以取categories[:-1] # 如果初始化OneHotEncoder时设置了dropLast=False,就直接取全部categories for cat in categories[:-1]: ohe_feature_names.append(f"{original_col}_{cat}") # 拼接数值类特征的名称 numeric_cols = ["Pclass", "Age", "Fare"] all_feature_names = ohe_feature_names + numeric_cols # 打印结果验证 print(all_feature_names)
输出结果格式和sklearn的get_feature_names_out完全对齐,针对泰坦尼克号数据集会输出类似如下内容:
['Sex_female', 'Sex_male', 'Embarked_S', 'Embarked_C', 'Embarked_Q', 'Pclass', 'Age', 'Fare']
注意:如果初始化StringIndexer时设置了
handleInvalid="keep",无效值会被映射到索引0位置,对应的类别名为__unknown,会自动出现在返回的类别列表中,不需要额外处理。
内容的提问来源于stack exchange,提问作者itscarlayall
相关产品推荐
相关产品推荐

