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

PySpark多列独热编码:统一展示所有分类列的唯一标签

PySpark多列分类特征独热编码:保留全局所有分类标签

要做到让每个分类列的独热编码包含所有分类列的全局唯一标签,得先统一全局标签集合,再基于这个集合对每个列编码,最后把向量拆成单独列。具体操作如下:

1. 收集全局所有分类标签

先从所有分类列里把所有唯一标签提取出来合并,确保每个列的编码都能覆盖这些标签:

from pyspark.sql import SparkSession
from pyspark.ml.feature import StringIndexer, OneHotEncoder, Pipeline
from pyspark.sql.functions import col, udf
from pyspark.sql.types import ArrayType, DoubleType

# 注意你原代码里cat_col1后面多了个空格,这里要去掉,不然会找不到列
categorical_columns = ['cat_col1', 'cat_col2']

# 提取所有分类列的全局唯一标签
all_categories = (
    temp_df.select(*categorical_columns)
    .rdd.flatMap(lambda row: row)
    .distinct()
    .collect()
)
all_categories.sort()  # 排序保证标签顺序一致,避免结果乱序

2. 构建统一编码的Pipeline

给每个分类列用全局标签集合做索引,再做独热编码:

# 每个列都用全局标签做索引,确保所有标签都被纳入
indexers = [
    StringIndexer(
        inputCol=c,
        outputCol=f"{c}_indexed",
        vocab=all_categories,
        handleInvalid="keep"  # 万一有未知标签,也能处理
    )
    for c in categorical_columns
]

# 独热编码时保留所有标签,不要丢弃最后一个类别
encoders = [
    OneHotEncoder(
        inputCol=indexer.getOutputCol(),
        outputCol=f"{c}_encoded",
        dropLast=False
    )
    for c, indexer in zip(categorical_columns, indexers)
]

# 运行Pipeline
pipeline = Pipeline(stages=indexers + encoders)
model = pipeline.fit(temp_df)
encoded_df = model.transform(temp_df)

3. 把独热编码向量拆成单独列

默认独热编码输出的是向量,得把它拆成一个个单独的列,列名对应「原列名_标签」:

# 定义UDF把向量转成数组
vector_to_array = udf(lambda vec: vec.toArray().tolist(), ArrayType(DoubleType()))

# 初始化结果表,先保留ID列
final_df = encoded_df.select("ID")

# 逐个处理每个分类列的编码结果
for c in categorical_columns:
    # 先把向量转成数组
    temp_df_array = encoded_df.withColumn(f"{c}_array", vector_to_array(col(f"{c}_encoded")))
    # 给每个标签生成对应的列
    for i, label in enumerate(all_categories):
        final_df = final_df.join(
            temp_df_array.select("ID", col(f"{c}_array").getItem(i).alias(f"{c}_{label}")),
            on="ID",
            how="inner"
        )

# 查看最终结果
final_df.display()

重点修正说明

  • 原代码里categorical_columns的cat_col1后面多了个空格,这会导致Spark找不到列,必须去掉。
  • 原代码的StringIndexer是针对单个列生成索引,只包含该列自己的标签,现在通过指定vocab=all_categories,让每个列的索引都覆盖全局所有标签。
  • 原代码没把独热编码的向量拆成单独列,所以展示的是索引列,现在通过UDF和数组拆分,生成了你要的展开式列。

内容的提问来源于stack exchange,提问作者AB95

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 09:35:28