如何用PySpark处理含动态Pin-State列的CSV文件并拆分行
PySpark 拆分多组Pin-State记录解决方案
问题描述
输入无表头CSV数据:
D,neel,32,1,pin1,state1,male D,sani,31,2,pin1,state1,pin2,state2,female D,raja,33,3,pin1,state1,pin2,state2,pin3,state3,male
需要转换为以下格式,核心规则是第4列的数字决定每条记录需要拆分出的Pin-State对数量:
D,neel,32,1,pin1,state1,male D,sani,31,2,pin1,state1,female D,sani,31,2,pin2,state2,female D,raja,33,3,pin1,state1,male D,raja,33,3,pin2,state2,male D,raja,33,3,pin3,state3,male
规则细节:
- 第4列数字为1时,保留原记录的1组Pin-State
- 第4列数字为2时,拆分成2条记录,每条对应一组Pin-State
- 第4列数字为3时,拆分成3条记录,每条对应一组Pin-State
解决方案代码
以下是完整的PySpark实现代码,直接运行即可得到目标输出:
from pyspark.sql import SparkSession # 初始化SparkSession spark = SparkSession.builder.appName("SplitPinStatePairs").getOrCreate() # 读取无表头CSV,按整行读取避免列数不一致报错 df = spark.read.text("input.csv") # 处理每行数据,拆分生成目标记录 def split_record(row): fields = row.value.split(",") # 提取固定前缀(前4列)和性别字段(最后1列) prefix = fields[:4] gender = fields[-1] pair_num = int(prefix[3]) # 按每2个字段一组拆分Pin-State对 pin_state_groups = [fields[4 + 2*i : 4 + 2*(i+1)] for i in range(pair_num)] # 拼接生成每条新记录 return [",".join(prefix + group + [gender]) for group in pin_state_groups] # 用RDD的flatMap实现一对多拆分,再转回DataFrame result_rdd = df.rdd.flatMap(split_record) result_df = result_rdd.toDF(["value"]) # 保存为无表头CSV,覆盖已有文件 result_df.write.mode("overwrite").option("header", "false").csv("output.csv") # 关闭SparkSession spark.stop()
代码说明
- 读取数据:用
read.text整行读取,避开因每行列数不同导致的读取失败问题 - 行处理逻辑:
- 拆分每行字段,分离固定前缀、性别字段和中间的Pin-State字段
- 根据第4列的数字,将Pin-State字段按每2个一组拆分
- 把前缀、单组Pin-State、性别拼接成新的完整记录
- RDD转换:通过
flatMap将单条原始记录拆分为多条目标记录 - 输出保存:将结果保存为无表头CSV,支持覆盖已有输出文件
纯DataFrame API替代实现
如果偏好只用DataFrame API,可结合自定义UDF实现:
from pyspark.sql import SparkSession from pyspark.sql.functions import udf, explode, array, col from pyspark.sql.types import ArrayType, StringType spark = SparkSession.builder.appName("SplitPinStatePairs").getOrCreate() # 按最大列数定义表头,读取CSV df = spark.read.csv( "input.csv", header=False, inferSchema=False, schema="_c0,_c1,_c2,_c3,_c4,_c5,_c6,_c7,_c8,_c9,_c10" ) # 定义UDF:拆分Pin-State对为列表 def split_pairs(pair_count, *cols): pairs = [] for i in range(int(pair_count)): pairs.append([cols[2*i], cols[2*i+1]]) return pairs split_udf = udf(split_pairs, ArrayType(ArrayType(StringType()))) # 应用UDF拆分记录,拼接成目标格式 result_df = df.withColumn( "pin_state_pairs", split_udf("_c3", *[f"_c{i}" for i in range(4, 11)]) ).withColumn("pair", explode("pin_state_pairs")).select( "_c0", "_c1", "_c2", "_c3", col("pair")[0].alias("pin"), col("pair")[1].alias("state"), # 根据原始记录长度动态选择性别列 col("_c10") if "_c10" in df.columns else col("_c8") if "_c8" in df.columns else col("_c6") ).withColumn( "value", array("_c0", "_c1", "_c2", "_c3", "pin", "state", "_c10").cast(StringType()) ) # 保存输出 result_df.write.mode("overwrite").option("header", "false").csv("output_df_api.csv") spark.stop()
内容的提问来源于stack exchange,提问作者neel banerjee
相关产品推荐
相关产品推荐

