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

基于AWS EMR优化10亿行PySpark特征工程性能的咨询

AWS EMR PySpark 10亿行交易数据特征工程性能优化方案

我在AWS EMR平台上用PySpark处理约10亿行的交易数据集,目标是完成特征工程:基于transaction_type划分不同交易分层,在多个历史时间窗口(如过去30天、90天)内,为每个客户计算sum、mean、standard deviation、entropy、Gini coefficient等统计特征,且需生成每日统计结果。

数据集Schema(transactions_df)

  • customer_id (String): 客户唯一ID
  • transaction_date (Date): 交易日期
  • transaction_amount (Double): 交易金额
  • transaction_type (String): 交易类型(如deposit、withdrawal等)

原简化代码

class FeatureEngineering:
    def __init__(self, transactions_df, strata_conditions):
        self.transactions_df = transactions_df
        self.strata_conditions = strata_conditions
        self.all_customers = transactions_df.select("customer_id").distinct()
        self.engineering_dates = (
            transactions_df.select("transaction_date").distinct().orderBy("transaction_date").collect()
        )
    
    def generate_strata_features(self, target_columns, time_periods):
        final_df = None

        for date in self.engineering_dates:  # Loop over each date
            for stratum_name, condition in self.strata_conditions.items():  # Loop over each stratum
                stratum_df = self.transactions_df.filter(condition)

                for col_name in target_columns:  # Loop over each target column
                    for period in time_periods:  # Loop over each time period
                        # Filter by the time period
                        filtered_df = self._filter_by_time_period(stratum_df, date, period)

                        # Compute statistics for the current time period and stratum
                        stats_df = self._compute_stats(filtered_df, stratum_name, period, col_name)
                        
                        # Join with the final dataframe
                        if final_df is None:
                            final_df = stats_df
                        else:
                            final_df = final_df.join(stats_df, on="customer_id", how="outer")

            # Add the date as a column to the final dataframe
            final_df = final_df.withColumn("engineering_date", F.lit(date["transaction_date"]))

        return final_df

    def _filter_by_time_period(self, df, date, period):
        """Filters the DataFrame based on the time period (e.g., last 30 days, last 90 days)."""
        return df.filter(
            (F.datediff(F.lit(date["transaction_date"]), F.col("transaction_date")) <= period) &
            (F.datediff(F.lit(date["transaction_date"]), F.col("transaction_date")) > 0)
        )
    
    def _compute_stats(self, df, stratum_name, period, col_name):
        """Computes stats like sum, mean, standard deviation, entropy, and Gini for a given stratum and period."""
        stats_df = df.groupBy("customer_id").agg(
            F.sum(col_name).alias(f"{stratum_name}_{period}_sum"),
            F.mean(col_name).alias(f"{stratum_name}_{period}_mean"),
            F.stddev(col_name).alias(f"{stratum_name}_{period}_stddev"),
            self._compute_entropy(df, col_name).alias(f"{stratum_name}_{period}_entropy"),
            self._compute_gini(df, col_name).alias(f"{stratum_name}_{period}_gini")
        )
        return stats_df
    
    def _compute_entropy(self, df, col_name):
        """Computes Shannon entropy for a given column."""
        total_count = df.groupBy('customer_id').agg(F.count(col_name).alias("total_count"))
        value_count = df.groupBy('customer_id', col_name).agg(F.count("*").alias("value_count"))
        df = value_count.join(total_count, on="customer_id").withColumn("prob", F.col("value_count") / F.col("total_count"))
        entropy_df = df.withColumn("entropy_component", -F.col("prob") * F.log2(F.col("prob")))
        return entropy_df.groupBy('customer_id').agg(F.sum("entropy_component").alias("entropy"))

    def _compute_gini(self, df, col_name):
        """Computes the Gini coefficient for a given column."""
        w = Window.partitionBy('customer_id').orderBy(col_name)
        df = df.withColumn("rank", F.sum(col_name).over(w))
        total_sum = df.groupBy('customer_id').agg(F.sum(col_name).alias("total_sum"))
        df = df.join(total_sum, on="customer_id").withColumn("norm_rank", F.col("rank") / F.col("total_sum"))
        df = df.withColumn("cum_norm_rank", F.sum("norm_rank").over(w))
        df = df.withColumn("gini", (2 * df.cum_norm_rank - 1) / df.count())
        return df.groupBy('customer_id').agg(F.mean("gini").alias("gini"))

当前代码性能问题

运行耗时极久(数小时),甚至写入S3前出现内存溢出,核心问题:

  1. 多层嵌套循环(遍历日期、分层、目标列、时间窗口),完全浪费PySpark分布式并行能力;
  2. 频繁Join操作触发大量Shuffle,导致内存不足;
  3. 重复扫描原始数据集,整体执行效率极低。

解决方案

1. 重构代码以利用PySpark并行性、避免内存问题

核心是把Driver端的循环逻辑转化为Spark分布式操作,让集群并行处理所有维度:

  • 取消Driver端日期循环:不要collect()日期列表,而是生成所有客户-日期的基准表,通过窗口函数计算每个日期下的历史窗口统计;
  • 分层逻辑提前打标签:给每条交易记录打上分层标签(用when函数匹配分层条件),后续按标签分组聚合,替代循环过滤;
  • 多时间窗口统一处理:定义多个窗口规格(对应不同时间周期),一次性计算所有窗口的特征,避免循环遍历时间窗口;
  • 合并统计计算:把sum、mean、熵、基尼系数的计算整合到同一批聚合/窗口操作中,避免重复扫描原始数据。

重构后核心代码示例:

from pyspark.sql import functions as F
from pyspark.sql.window import Window
from pyspark.sql.types import StorageLevel

class OptimizedFeatureEngineering:
    def __init__(self, transactions_df, strata_conditions):
        self.transactions_df = transactions_df.repartition("customer_id")  # 提前按客户分区,减少后续Shuffle
        # 给每条交易打分层标签
        self.tagged_transactions = self._tag_strata(strata_conditions).persist(StorageLevel.MEMORY_AND_DISK_SER)
        # 生成客户-日期基准表(确保每个客户每日都有记录)
        self.customer_date_base = self.transactions_df.groupBy("customer_id", "transaction_date") \
            .count().drop("count").withColumnRenamed("transaction_date", "engineering_date").persist(StorageLevel.MEMORY_AND_DISK_SER)

    def _tag_strata(self, strata_conditions):
        df = self.transactions_df
        # 给每个分层打标记列,非互斥分层适用
        for stratum_name, condition in strata_conditions.items():
            df = df.withColumn(stratum_name, F.when(condition, 1).otherwise(0))
        return df

    def generate_all_features(self, target_columns, time_periods):
        # 关联基准表与交易数据
        joined_df = self.customer_date_base.join(
            self.tagged_transactions,
            on=["customer_id"],
            how="left"
        )

        final_df = self.customer_date_base
        for period in time_periods:
            # 定义滑动时间窗口:过去period天(不含当天)
            window_spec = Window.partitionBy("customer_id", "engineering_date") \
                .orderBy(F.col("transaction_date").cast("timestamp").cast("long")) \
                .rangeBetween(-period*86400, -1)

            # 计算基础统计量
            stats_df = joined_df
            for col in target_columns:
                stats_df = stats_df \
                    .withColumn(f"{col}_{period}_sum", F.sum(col).over(window_spec)) \
                    .withColumn(f"{col}_{period}_mean", F.avg(col).over(window_spec)) \
                    .withColumn(f"{col}_{period}_stddev", F.stddev(col).over(window_spec))

            # 计算熵和基尼系数
            stats_df = self._compute_entropy_window(stats_df, period, window_spec, target_columns[0])
            stats_df = self._compute_gini_window(stats_df, period, window_spec, target_columns[0])

            # 按分层提取特征并合并
            for stratum_name in strata_conditions.keys():
                stratum_features = stats_df.filter(F.col(stratum_name) == 1) \
                    .select(
                        "customer_id", "engineering_date",
                        *[F.col(f"{col}_{period}_{stat}").alias(f"{stratum_name}_{period}_{stat}") 
                          for col in target_columns for stat in ["sum", "mean", "stddev", "entropy", "gini"]]
                    )
                final_df = final_df.join(stratum_features, on=["customer_id", "engineering_date"], how="left")

        return final_df.fillna(0)  # 填充无交易的客户特征为0

    def _compute_entropy_window(self, df, period, window_spec, col_name):
        # 窗口内计算金额分布的熵
        df = df.withColumn(f"{col_name}_{period}_total", F.count(col_name).over(window_spec)) \
            .withColumn(f"{col_name}_{period}_val_count", F.count(col_name).over(
                Window.partitionBy("customer_id", "engineering_date", col_name)
            )) \
            .withColumn(f"{col_name}_{period}_prob", F.col(f"{col_name}_{period}_val_count") / F.col(f"{col_name}_{period}_total")) \
            .withColumn(f"{col_name}_{period}_entropy", F.sum(-F.col(f"{col_name}_{period}_prob") * F.log2(F.col(f"{col_name}_{period}_prob"))).over(window_spec))
        return df.drop(f"{col_name}_{period}_total", f"{col_name}_{period}_val_count", f"{col_name}_{period}_prob")

    def _compute_gini_window(self, df, period, window_spec, col_name):
        # 窗口内计算基尼系数
        sorted_window = window_spec.orderBy(col_name)
        df = df.withColumn(f"{col_name}_{period}_cum_sum", F.sum(col_name).over(sorted_window)) \
            .withColumn(f"{col_name}_{period}_total_sum", F.sum(col_name).over(window_spec)) \
            .withColumn(f"{col_name}_{period}_cum_norm", F.col(f"{col_name}_{period}_cum_sum") / F.col(f"{col_name}_{period}_total_sum")) \
            .withColumn(f"{col_name}_{period}_gini", (2 * F.sum(f"{col_name}_{period}_cum_norm").over(sorted_window) - 1) / F.count(col_name).over(window_spec))
        return df.drop(f"{col_name}_{period}_cum_sum", f"{col_name}_{period}_total_sum", f"{col_name}_{period}_cum_norm")

2. 缓存、Checkpointing与减少Shuffle的最佳实践

  • 缓存策略:
    • 只缓存重复使用的中间DataFrame,比如tagged_transactions和customer_date_base,用persist(StorageLevel.MEMORY_AND_DISK_SER)(序列化存储减少内存占用);
    • 缓存后记得在任务结束时调用unpersist()释放资源。
  • Checkpointing:
    • 当DAG依赖链过长时,将关键中间结果写入S3/HDFS做Checkpoint,切断DAG依赖,避免重复计算:
      spark.sparkContext.setCheckpointDir("s3://your-bucket/checkpoint/")
      intermediate_df = stats_df.checkpoint()
      
  • 减少Shuffle操作:
    • 优先用窗口函数替代groupBy+join,窗口函数仅按partitionBy列做一次Shuffle;
    • 提前按customer_id分区原始数据集,后续所有客户维度的操作无需再Shuffle;
    • 用broadcast广播小表:如果分层规则对应的是小数据集,用F.broadcast()将其广播到所有节点;
    • 避免不必要的distinct:生成客户-日期基准表时,用groupBy替代distinct+crossJoin,减少数据扫描。

3. 减少/消除嵌套循环的方法

  • 将循环维度转化为列:
    • 分层维度:给每条交易打分层标签,把分层过滤转化为列分组,替代循环遍历分层;
    • 时间窗口:定义多个窗口规格,一次性计算所有窗口的特征,或用array+explode生成窗口参数列,通过一次计算处理所有窗口;
  • 避免Driver端遍历数据:
    • 不要collect()日期、客户等数据集,而是通过Spark分布式操作生成基准表,替代Driver端的循环;
    • 多目标列处理:用动态生成聚合表达式的方式,一次性处理所有目标列的统计,避免循环遍历目标列。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 06:08:09