You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

大数据集Pandas慢循环优化问询:多列处理提速与PySpark适配

问题描述

我有一个包含月度数据的Pandas DataFrame(df),初始有a、b、c、d四列,当前代码执行以下三步操作:

  1. 生成所有无重复的列组合,并计算每组组合的求和列
  2. 将DataFrame扩展5个月的时间范围
  3. 为每一列生成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(),减少重复计算

三、整体流程核心优化建议

  1. 砍掉不必要的列组合:2万列的全组合完全不可行,需根据业务规则过滤组合(比如只保留特定维度的组合)
  2. 延迟计算策略:尽量在数据使用环节再生成移位列或组合列,而非提前全量存储
  3. 数据类型压缩:将数值列设置为更小的类型(比如float32代替float64),降低内存占用

内容的提问来源于stack exchange,提问作者jack homareau

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.24 05:30:10