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

如何获取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 23:40:36