PySpark实现将含列表的列转换为布尔列
PySpark实现逗号分隔标签列转二进制特征列
解决方案步骤
- 将Z列的逗号分隔字符串转换为数组类型
- 提取所有唯一标签值(作为新列名)
- 为每个标签生成二进制列,判断该行是否包含对应标签
完整代码示例
from pyspark.sql import SparkSession from pyspark.sql.functions import split, array_contains, when, col # 初始化Spark会话 spark = SparkSession.builder.appName("TagToBinaryCols").getOrCreate() # 创建示例DataFrame data = [ (1, 1, 1, "one,two,three"), (2, 1, 2, "one,two,four,five"), (3, 2, 1, "four,five") ] df = spark.createDataFrame(data, ["Id", "X", "Y", "Z"]) # 将Z列的逗号分隔字符串转为数组 df_with_array = df.withColumn("z_array", split(col("Z"), ",")) # 提取所有唯一标签(仅收集标签数据,避免拉取全量行导致集群崩溃) unique_tags = df_with_array.select("z_array") \ .rdd.flatMap(lambda row: row[0]) \ .distinct() \ .collect() # 为每个标签生成二进制列:存在则为1,否则为0 for tag in unique_tags: df_with_array = df_with_array.withColumn( tag, when(array_contains(col("z_array"), tag), 1).otherwise(0) ) # 整理最终结果,移除中间列 final_df = df_with_array.drop("Z", "z_array") final_df.show()
关键说明
- 避免全量数据collect:仅提取唯一标签值,数据量远小于全量行,不会压垮Driver节点
- 无需explode拆分:直接通过
array_contains判断数组内是否存在标签,避免行膨胀后再聚合的性能损耗
内容的提问来源于stack exchange,提问作者Fredrik Haarde
相关产品推荐
相关产品推荐

