自定义Target Mean Encoding后PySpark DataFrame.show()触发OOM错误求助
问题描述
我用PySpark读取了一个包含2400万行、约14个特征的CSV数据集,预处理前调用.show()能正常运行。但对其中一列应用自定义目标均值编码(Mean Encoding)函数后,调用该列的.show()时触发了OOM(内存溢出)错误。我已经尝试调整spark.driver.memory,但问题依然存在,想知道这是数据集损坏还是有其他未知问题?
原代码
from neo4j import GraphDatabase import pandas as pd import numpy as np import networkx as nx import xgboost as xgb from pyspark.sql import SparkSession import pyspark.sql.functions as F from pyspark.sql.functions import split, explode, col, avg, expr, create_map from pyspark.ml.feature import OneHotEncoder, StringIndexer def target_mean_encoding(df, col, target): """ :param df: pyspark.sql.dataframe dataframe to apply target mean encoding :param col: str list list of columns to apply target encoding :param target: str target column :return: dataframe with target encoded columns """ target_encoded_columns_list = [] for c in col: indexSuffix = '_indexed' indexer = StringIndexer( inputCol=c, outputCol=c+indexSuffix, handleInvalid='keep') indexer.setHandleInvalid('keep') model = indexer.fit(df) df = model.transform(df) df = df.drop(c) df = df.withColumnRenamed(c+indexSuffix, c) means = df.groupby(F.col(c)).agg(F.mean(target).alias(f"{c}_mean_encoding")) dict_ = means.toPandas().to_dict() target_encoded_columns = [F.when(F.col(c) == v, encoder) for v, encoder in zip(dict_[c].values(), dict_[f"{c}_mean_encoding"].values())] target_encoded_columns_list.append(F.coalesce(*target_encoded_columns).alias(f"{c}_mean_encoding")) df.select(*target_encoded_columns_list).show(1) df.show() return df.select( *target_encoded_columns_list) spark = SparkSession.builder.master( 'local[*]').config("spark.driver.memory", "10g").appName('sl-app').getOrCreate() df = spark.read.csv("data.csv", header=True, escape=",") indexer = StringIndexer(inputCol="Is Fraud?", outputCol="Is Fraud?_LE").setHandleInvalid("keep") indexed = indexer.fit(df).transform(df) indexed = indexed.drop('Is Fraud?') df = indexed.withColumnRenamed('Is Fraud?_LE', 'Is Fraud?') df_target_encoded = target_mean_encoding( df, col=['Merchant City'], target='Is Fraud?')
报错信息
C:\Users\Hai Nguyen\AppData\Local\Programs\Python\Python310\lib\site-packages\numpy\.libs\libopenblas.FB5AE2TYXYH2IJRDKGDGQ3XBKLKTF43H.gfortran-win_amd64.dll warnings.warn("loaded more than 1 DLL from .libs:") Picked up _JAVA_OPTIONS: -Xmx512M Picked up _JAVA_OPTIONS: -Xmx512M Setting default log level to "WARN". To adjust logging level use sc.setLogLevel(newLevel). For SparkR, use setLogLevel(newLevel). 23/04/19 15:10:53 WARN NativeCodeLoader: Unable to load native-hadoop library for your platform... using builtin-java classes where applicable 23/04/19 15:11:58 WARN package: Truncated the string representation of a plan since it was too large. This behavior can be adjusted by setting 'spark.sql.debug.maxToStringFields'. Exception in thread "refresh progress" Exception in thread "RemoteBlock-temp-file-clean-thread" java.lang.OutOfMemoryError: Java heap space java.lang.OutOfMemoryError: Java heap space at java.base/java.lang.invoke.DirectMethodHandle.allocateInstance(DirectMethodHandle.java:520) at java.base/java.lang.invoke.DirectMethodHandle$Holder.newInvokeSpecial(DirectMethodHandle$Holder) at java.base/java.lang.invoke.Invokers$Holder.linkToTargetMethod(Invokers$Holder) at org.apache.spark.storage.BlockManager$RemoteBlockDownloadFileManager.org$apache$spark$storage$BlockManager$RemoteBlockDownloadFileManager$$keepCleaning(BlockManager.scala:2158) at org.apache.spark.storage.BlockManager$RemoteBlockDownloadFileManager$$anon$2.run(BlockManager.scala:2124) 23/04/19 15:12:08 ERROR Utils: uncaught error in thread Spark Context Cleaner, stopping SparkContext java.lang.OutOfMemoryError: Java heap space 23/04/19 15:12:08 ERROR Utils: throw uncaught fatal error in thread Spark Context Cleaner java.lang.OutOfMemoryError: Java heap space Exception in thread "Spark Context Cleaner" java.lang.OutOfMemoryError: Java heap space Traceback (most recent call last): File "c:\Users\Hai Nguyen\Desktop\FPT\data2\tempCodeRunnerFile.py", line 142, in <module> df_target_encoded = target_mean_encoding( File "c:\Users\Hai Nguyen\Desktop\FPT\data2\tempCodeRunnerFile.py", line 92, in target_mean_encoding df.select(*target_encoded_columns_list).show(1) File "C:\Users\Hai Nguyen\AppData\Local\Programs\Python\Python310\lib\site-packages\pyspark\sql\dataframe.py", line 899, in show print(self._jdf.showString(n, 20, vertical)) File "C:\Users\Hai Nguyen\AppData\Local\Programs\Python\Python310\lib\site-packages\py4j\java_gateway.py", line 1322, in __call__ return_value = get_return_value( File "C:\Users\Hai Nguyen\AppData\Local\Programs\Python\Python310\lib\site-packages\pyspark\errors\exceptions\captured.py", line 169, in deco return f(*a, **kw) File "C:\Users\Hai Nguyen\AppData\Local\Programs\Python\Python310\lib\site-packages\py4j\protocol.py", line 326, in get_return_value raise Py4JJavaError( py4j.protocol.Py4JJavaError: An error occurred while calling o40441.showString. : java.lang.OutOfMemoryError: Java heap space
问题原因分析
- Pandas转换导致Driver内存过载:代码中
means.toPandas().to_dict()将分组后的结果全量加载到Driver端内存。如果Merchant City的基数很高(比如几十万个不同取值),Driver内存会被直接撑爆,这是OOM的核心原因。 - 冗余
.show()调用放大压力:循环内的df.show()会触发全量2400万行数据的计算和拉取,本地模式下所有数据都会集中到Driver端,进一步加剧内存不足。 - 环境变量覆盖Spark内存配置:报错中显示
Picked up _JAVA_OPTIONS: -Xmx512M,这个系统级Java堆内存设置会覆盖你配置的spark.driver.memory=10g,导致Driver实际可用内存只有512M,这也是调整内存无效的关键原因。
修复方案
1. 用Spark原生关联替代Pandas字典映射
放弃将分组结果转成Pandas字典的方式,改用Spark的join操作实现均值编码映射,全程分布式处理,避免Driver内存过载。同时使用F.broadcast()广播小数据集,提升关联效率。
2. 移除或限制.show()调用
调试时用.limit(N).show()替代全量show(),避免拉取所有数据到Driver端;非必要情况下直接移除循环内的show()调用。
3. 修正Java内存环境变量
删除系统环境变量_JAVA_OPTIONS中的-Xmx512M配置,确保Spark的内存参数生效。
修复后的代码
from pyspark.sql import SparkSession import pyspark.sql.functions as F from pyspark.ml.feature import StringIndexer def target_mean_encoding(df, cols, target): """ :param df: pyspark.sql.dataframe 要应用目标均值编码的数据集 :param cols: str list 需要编码的列名列表 :param target: str 目标列名 :return: 带目标编码列的数据集 """ result_df = df for c in cols: # 字符串索引转换(如果原列是字符串类型,转索引可提升后续分组效率) indexSuffix = '_indexed' indexer = StringIndexer(inputCol=c, outputCol=c+indexSuffix, handleInvalid='keep') model = indexer.fit(result_df) result_df = model.transform(result_df).drop(c).withColumnRenamed(c+indexSuffix, c) # 计算每个分组的目标均值 mean_df = result_df.groupBy(c).agg(F.mean(target).alias(f"{c}_mean_encoding")) # 广播小数据集后关联,替代字典映射 result_df = result_df.join(F.broadcast(mean_df), on=c, how='left') return result_df # 初始化SparkSession,补充executor内存配置 spark = SparkSession.builder.master('local[*]')\ .config("spark.driver.memory", "10g")\ .config("spark.executor.memory", "8g")\ .appName('sl-app').getOrCreate() # 读取数据集 df = spark.read.csv("data.csv", header=True, escape=",") # 目标列字符串索引转换 indexer = StringIndexer(inputCol="Is Fraud?", outputCol="Is Fraud?_LE", handleInvalid="keep") df = indexer.fit(df).transform(df).drop('Is Fraud?').withColumnRenamed('Is Fraud?_LE', 'Is Fraud?') # 应用目标均值编码 df_target_encoded = target_mean_encoding(df, cols=['Merchant City'], target='Is Fraud?') # 调试时仅查看前10行数据 df_target_encoded.select('Merchant City', 'Merchant City_mean_encoding').limit(10).show()
内容的提问来源于stack exchange,提问作者Hai Nguyen
相关产品推荐
相关产品推荐

