PySpark向DataFrame新增列出错:列值重复问题排查与解决
PySpark合并股票Close列重复值问题排查与最优解法
问题重现
用户尝试将不同股票的Close列合并到同一PySpark DataFrame中,编写代码如下:
small_list = ["INFY","TCS", "SBIN", "ICICIBANK"] frame = spark_frame.where(col("symbol") == small_list[0]).select('close') ## spark frame is a pyspark.sql.dataframe.DataFrame for single_stock in small_list[1:]: print(single_stock) current_stock = spark_frame.where(col("symbol") == single_stock).select(['close']) current_stock.collect() frame.collect() frame = frame.withColumn(single_stock, current_stock.close)
执行后发现所有新增列的值与第一列Close完全重复,结果示例:
[Row(close=736.85, TCS=736.85, SBIN=736.85, ICICIBANK=736.85), Row(close=734.7, TCS=734.7, SBIN=734.7, ICICIBANK=734.7), Row(close=746.0, TCS=746.0, SBIN=746.0, ICICIBANK=746.0), Row(close=738.85, TCS=738.85, SBIN=738.85, ICICIBANK=738.85)]
原spark_frame结构示例:
[Row(SYMBOL='LINC', SERIES=' EQ', TIMESTAMP=datetime.datetime(2021, 12, 20, 0, 0), PREVCLOSE=235.6, OPEN=233.95, HIGH=234.0, LOW=222.15, LAST=222.15, CLOSE=224.2, AVG_PRICE=226.63, TOTTRDQTY=6447, TOTTRDVAL=14.61, TOTALTRADES=206, DELIVQTY=5507, DELIVPER=85.42), Row(SYMBOL='LINC', SERIES=' EQ', TIMESTAMP=datetime.datetime(2021, 12, 21, 0, 0), PREVCLOSE=224.2, OPEN=243.85, HIGH=243.85, LOW=222.85, LAST=226.0, CLOSE=225.6, AVG_PRICE=227.0, TOTTRDQTY=8447, TOTTRDVAL=19.17, TOTALTRADES=266, DELIVQTY=3401, DELIVPER=40.26), Row(SYMBOL='SCHAEFFLER', SERIES=' EQ', TIMESTAMP=datetime.datetime(2020, 8, 6, 0, 0), PREVCLOSE=3593.9, OPEN=3611.85, HIGH=3618.35, LOW=3542.5, LAST=3594.95, CLOSE=3573.1, AVG_PRICE=3580.73, TOTTRDQTY=12851, TOTTRDVAL=460.16, TOTALTRADES=1886, DELIVQTY=9649, DELIVPER=75.08), Row(SYMBOL='SCHAEFFLER', SERIES=' EQ', TIMESTAMP=datetime.datetime(2020, 8, 7, 0, 0), PREVCLOSE=3573.1, OPEN=3591.0, HIGH=3591.0, LOW=3520.0, LAST=3548.95, CLOSE=3543.85, AVG_PRICE=3554.6, TOTTRDQTY=2406, TOTTRDVAL=85.52, TOTALTRADES=688, DELIVQTY=1452, DELIVPER=60.35)]
期望得到按时间对齐的宽格式DataFrame:
[Row(LINC=224.2, SCHAEFFLER=3573.1), Row(LINC=225.6, SCHAEFFLER=3543.85)]
错误原因
- 无关联的列拼接:使用
frame.withColumn(single_stock, current_stock.close)时,未指定两个DataFrame的关联条件(如TIMESTAMP),PySpark会将current_stock的close列视为常量广播到frame的每一行,导致所有列值重复。 - 冗余的
collect()调用:current_stock.collect()和frame.collect()完全多余,这两个方法会把分布式数据拉到Driver节点,既浪费资源又破坏PySpark的分布式计算特性,对结果无任何帮助。
最优解决方法
PySpark中处理长格式转宽格式(将不同分组的列展开)最高效的方式是使用**pivot**函数,按时间戳分组后透视股票代码,提取对应Close值:
代码实现
from pyspark.sql import functions as F # 按TIMESTAMP分组,透视SYMBOL列,取每个分组下的CLOSE值 result_df = spark_frame.groupBy("TIMESTAMP")\ .pivot("SYMBOL")\ .agg(F.first("CLOSE"))\ .drop("TIMESTAMP") # 不需要时间戳列可删除 # 若仅需指定股票列表,在pivot中传入参数过滤 small_list = ["INFY","TCS", "SBIN", "ICICIBANK"] result_df = spark_frame.groupBy("TIMESTAMP")\ .pivot("SYMBOL", small_list)\ .agg(F.first("CLOSE"))\ .drop("TIMESTAMP") result_df.show()
代码说明
groupBy("TIMESTAMP"):按时间戳分组,确保同一时间点的股票数据被放在同一行。pivot("SYMBOL"):将SYMBOL列的不同值转换为新的列名;指定small_list作为第二个参数时,仅保留列表中的股票列,过滤无关数据。agg(F.first("CLOSE")):每个分组下取唯一的Close值(假设同一股票同一时间只有一条数据),也可使用F.max或F.min,效果一致。drop("TIMESTAMP"):若最终结果不需要时间戳列,直接删除即可得到期望结构。
这种方法比循环拼接高效得多,充分利用PySpark的分布式计算能力,避免了循环带来的性能问题和逻辑错误。
内容的提问来源于stack exchange,提问作者Inder
相关产品推荐
相关产品推荐

