基于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前出现内存溢出,核心问题:
- 多层嵌套循环(遍历日期、分层、目标列、时间窗口),完全浪费PySpark分布式并行能力;
- 频繁Join操作触发大量Shuffle,导致内存不足;
- 重复扫描原始数据集,整体执行效率极低。
解决方案
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()释放资源。
- 只缓存重复使用的中间DataFrame,比如
- Checkpointing:
- 当DAG依赖链过长时,将关键中间结果写入S3/HDFS做Checkpoint,切断DAG依赖,避免重复计算:
spark.sparkContext.setCheckpointDir("s3://your-bucket/checkpoint/") intermediate_df = stats_df.checkpoint()
- 当DAG依赖链过长时,将关键中间结果写入S3/HDFS做Checkpoint,切断DAG依赖,避免重复计算:
- 减少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
相关产品推荐
相关产品推荐

