PySpark中遍历指定列匹配值并生成目标DataFrame的方法
基于Spark实现指定规则的DataFrame转换
初始数据
初始DataFrame结构及数据如下:
| a | b | id | m2000 | m2001 | m2002 | m2003 | m2004 | m2005 |
|---|---|---|---|---|---|---|---|---|
| a | world | 1 | 0 | 0 | 1 | 0 | 0 | 1 |
需求说明
需要从m2000至m2014的列中筛选出值为1的列,生成新的DataFrame,规则如下:
- 保留原表的
id列 year列为固定前缀10/10/拼接值为1的列名中的年份(示例中为2002)yearend列为固定前缀12/12/拼接值为1的列名中的年份(示例中为2005)- 每个值为1的列对应生成一行数据(示例中因有两个符合条件的列,生成两行相同数据)
初始DataFrame创建代码
from pyspark.shell import spark from pyspark.sql.types import StructType, StructField, StringType, IntegerType data2 = [("a", "world", "1", 0, 0, 1, 0, 0, 1),] schema = StructType([ \ StructField("a", StringType(), True), \ StructField("b", StringType(), True), \ StructField("id", StringType(), True), \ StructField("m2000", IntegerType(), True), \ StructField("m2001", IntegerType(), True), \ StructField("m2002", IntegerType(), True), \ StructField("m2003", IntegerType(), True), \ StructField("m2004", IntegerType(), True), \ StructField("m2005", IntegerType(), True), \ ]) df = spark.createDataFrame(data=data2, schema=schema) df.printSchema() df.show(truncate=False)
解决方案代码
from pyspark.sql.functions import lit, regexp_extract, when, array, array_filter, explode, first, last from pyspark.sql.window import Window # 1. 筛选出m2000到m2014的列名 year_cols = [col for col in df.columns if col.startswith('m') and 2000 <= int(col[1:]) <= 2014] # 2. 构造数组,收集每行中值为1的列对应的年份,过滤掉空值 year_array = array(*[ when(col(c) == 1, regexp_extract(lit(c), r'm(\d{4})', 1)).otherwise(None) for c in year_cols ]) filtered_years = array_filter(year_array, lambda x: x.isNotNull()) # 3. 展开数组,每个年份对应一行数据 df_expanded = df.withColumn('year_str', explode(filtered_years)) # 4. 按id分组,获取第一个和最后一个符合条件的年份 window_spec = Window.partitionBy('id') df_with_bounds = df_expanded.withColumn('first_year', first('year_str').over(window_spec)) \ .withColumn('last_year', last('year_str').over(window_spec)) # 5. 构造最终的year和yearend列,生成结果DataFrame result_df = df_with_bounds.select( 'id', (lit('10/10/') + df_with_bounds.first_year).alias('year'), (lit('12/12/') + df_with_bounds.last_year).alias('yearend') ) # 查看结果 result_df.show(truncate=False)
执行结果
运行上述代码后,得到的结果DataFrame如下:
+---+------------+------------+ |id |year |yearend | +---+------------+------------+ |1 |10/10/2002 |12/12/2005 | |1 |10/10/2002 |12/12/2005 | +---+------------+------------+
内容的提问来源于stack exchange,提问作者lunbox
相关产品推荐
相关产品推荐

