大数据集Pandas慢循环优化问询:多列处理提速与PySpark适配
问题描述
我有一个包含月度数据的Pandas DataFrame(df),初始有a、b、c、d四列,当前代码执行以下三步操作:
- 生成所有无重复的列组合,并计算每组组合的求和列
- 将DataFrame扩展5个月的时间范围
- 为每一列生成1-5个月的移位后缀列
但当初始列数达到20000时,在Azure Databricks环境中第三步运行极慢。现咨询:
- 如何优化第三步或整体代码
- 该场景是否更适合使用PySpark(后续还会持续增加更多列)
原代码如下:
import pandas as pd import itertools as it pd.set_option('display.max_rows', None) pd.set_option('display.max_columns', None) df = pd.DataFrame({'date':pd.to_datetime(['01-31-2023','02-28-2023','03-31-2023','04-30-2023', '05-31-2023']), 'a': [3,4,5,6,3], 'b': [5,7,1,0,5], 'c':[3,4,2,1,3], 'd':[2,0,1,5,9]}).set_index('date') #1. 生成所有列组合的求和列 orig_cols = df.columns for r in range(2, df.shape[1] + 1): for cols in it.combinations(orig_cols, r): df["_".join(cols)] = df.loc[:, cols].sum(axis=1) #2. 扩展DataFrame至5个月后 first_date = df.first_valid_index() end_forecast_horizon = '2023-10-30' min_lead = 1 max_lead = 5 expand_dates = pd.date_range(first_date, end_forecast_horizon, freq='M') df = df.reindex(expand_dates) #3. 为每列生成1-5个月的移位列 for col in df.columns: for i in range(min_lead,max_lead): df["%s_%s"%(col,i)] = df[col].shift(i) df
优化方案
一、Pandas环境下的代码优化
1. 第三步移位操作的批量优化
原代码循环每一列生成移位列,当列数达2万时,循环次数会达到8万次,每次新增列都会触发DataFrame内存重分配,这是核心性能瓶颈。改用批量移位+合并的方式,大幅减少内存操作:
# 替换第三步的循环代码 shifted_dfs = [] for i in range(min_lead, max_lead): # 对所有列同时移位并重命名 shifted_df = df.shift(i).add_suffix(f"_{i}") shifted_dfs.append(shifted_df) # 一次性合并所有移位后的结果 df = pd.concat([df] + shifted_dfs, axis=1)
2. 第一步列组合生成的致命问题优化
当初始列数为2万时,全列组合的数量是2^20000 - 20001,这会直接导致内存溢出,比第三步的问题更严重,必须调整:
- 若组合列用于建模,考虑延迟计算:仅在需要使用时生成对应组合,而非提前全量存储
- 若必须生成,使用稀疏矩阵(Pandas的
SparseDataFrame或scipy.sparse)存储,利用组合列的重复/缺失特性压缩内存 - 复用已生成的组合列:比如
a_b_c的和等于a_b + c,避免重复计算
二、切换到PySpark的必要性与适配建议
当列数持续增长到万级以上时,PySpark是更合适的选择:
- 基于分布式计算框架,可横向扩展集群资源处理大规模数据
- 列操作矢量化,底层优化了内存管理,避免Pandas单进程内存瓶颈
PySpark适配示例代码
from pyspark.sql import SparkSession from pyspark.sql import functions as F from itertools import combinations spark = SparkSession.builder.appName("LargeColumns").getOrCreate() # 1. 读取初始数据 df = spark.createDataFrame( [ ("2023-01-31", 3,5,3,2), ("2023-02-28",4,7,4,0), ("2023-03-31",5,1,2,1), ("2023-04-30",6,0,1,5), ("2023-05-31",3,5,3,9) ], schema=["date", "a", "b", "c", "d"] ).withColumn("date", F.to_date("date", "yyyy-MM-dd")).orderBy("date") # 2. 生成列组合求和列 orig_cols = df.columns[1:] for r in range(2, len(orig_cols)+1): for cols in combinations(orig_cols, r): col_name = "_".join(cols) df = df.withColumn(col_name, sum(F.col(c) for c in cols)) # 3. 扩展时间范围 min_date = df.select(F.min("date")).first()[0] date_df = spark.createDataFrame( pd.date_range(min_date, "2023-10-30", freq="M").to_frame(name="date"), schema=["date"] ) df = date_df.join(df, on="date", how="left") # 4. 生成移位列 min_lead = 1 max_lead = 5 df = df.orderBy("date") for i in range(min_lead, max_lead): window_spec = F.window.orderBy("date").rowsBetween(-i, -i) # 批量生成移位列表达式,减少withColumn调用次数 shift_cols = [F.first(col).over(window_spec).alias(f"{col}_{i}") for col in df.columns[1:]] df = df.select("*", *shift_cols)
PySpark优化要点
- 用列表推导式批量生成列表达式,减少
withColumn的重复调用 - 根据集群资源调整DataFrame分区数,避免分区过多/过少
- 对频繁使用的中间结果调用
df.cache(),减少重复计算
三、整体流程核心优化建议
- 砍掉不必要的列组合:2万列的全组合完全不可行,需根据业务规则过滤组合(比如只保留特定维度的组合)
- 延迟计算策略:尽量在数据使用环节再生成移位列或组合列,而非提前全量存储
- 数据类型压缩:将数值列设置为更小的类型(比如
float32代替float64),降低内存占用
内容的提问来源于stack exchange,提问作者jack homareau
相关产品推荐
相关产品推荐

