使用PySpark Bucketizer时为何报IllegalArgumentException:splits参数无效?
PySpark Bucketizer IllegalArgumentException 问题解决
问题描述
使用PySpark的Bucketizer进行分箱时触发IllegalArgumentException,提示splits参数无效,相关代码及报错信息如下:
原代码
import pyspark.sql.functions as F from pyspark.ml.feature import Bucketizer df = spark.createDataFrame( [{'score': 0.056906660760916945}, {'score': 0.014312104993006614}, {'score': 0.019505725666714737}, {'score': 0.05695818453036461}, {'score': 0.004712143467991172}, {'score': 0.01950581558579922}, {'score': 0.0}, {'score': 0.004479459183469148}, {'score': 0.0}, {'score': 0.0002634215537602458}, {'score': 1.6}, {'score': 0.0002634215537602458}], ["score"] ) splits = [0, 0.7, 0.21, 1] bucket_df = Bucketizer( splits=splits, inputCol="score", outputCol="score_bucket" ).transform(df.where(F.col("score").isNotNull())) bucket_df.groupBy("score_bucket").count().show()
报错信息
--------------------------------------------------------------------------- IllegalArgumentException Traceback (most recent call last) <ipython-input-94-19c8c78fe553> in <cell line: 15>() 15 bucket_df = Bucketizer( 16 splits=splits, inputCol="score", outputCol="score_bucket" ---> 17 ).transform(df.where(F.col("score").isNotNull())) 18 bucket_df.groupBy("score_bucket").count().show() ~/Downloads/2023-06-01-spark/spark-3.4.0-bin-hadoop3/python/pyspark/ml/base.py in transform(self, dataset, params) 260 return self.copy(params)._transform(dataset) 261 else: ---> 262 return self._transform(dataset) 263 else: 264 raise TypeError("Params must be a param map but got %s." % type(params)) ~/Downloads/2023-06-01-spark/spark-3.4.0-bin-hadoop3/python/pyspark/ml/wrapper.py in _transform(self, dataset) 395 assert self._java_obj is not None 396 ---> 397 self._transfer_params_to_java() 398 return DataFrame(self._java_obj.transform(dataset._jdf), dataset.sparkSession) 399 ~/Downloads/2023-06-01-spark/spark-3.4.0-bin-hadoop3/python/pyspark/ml/wrapper.py in _transfer_params_to_java(self) 169 for param in self.params: 170 if self.isSet(param): ---> 171 pair = self._make_java_param_pair(param, self._paramMap[param]) 172 self._java_obj.set(pair) 173 if self.hasDefault(param): ~/Downloads/2023-06-01-spark/spark-3.4.0-bin-hadoop3/python/pyspark/ml/wrapper.py in _make_java_param_pair(self, param, value) 158 java_param = self._java_obj.getParam(param.name) 159 java_value = _py2java(sc, value) ---> 160 return java_param.w(java_value) 161 162 def _transfer_params_to_java(self) -> None: ~/.python_venvs/pandars310/lib/python3.10/site-packages/py4j/java_gateway.py in __call__(self, *args) 1319 1320 answer = self.gateway_client.send_command(command) -> 1321 return_value = get_return_value( 1322 answer, self.gateway_client, self.target_id, self.name) 1323 ~/Downloads/2023-06-01-spark/spark-3.4.0-bin-hadoop3/python/pyspark/errors/exceptions/captured.py in deco(*a, **kw) 173 # Hide where the exception came from that shows a non-Pythonic 174 # JVM exception message. ---> 175 raise converted from None 176 else: 177 raise IllegalArgumentException: Bucketizer_aff53d93ca5b parameter splits given invalid value [0.0,0.7,0.21,1.0].
错误原因
Bucketizer对splits参数有强制约束:必须是严格递增的数值序列。原代码中splits = [0, 0.7, 0.21, 1]存在递减值(0.7 > 0.21),违反参数规则导致报错。此外,原splits未覆盖数据中的边界外值(如score=1.6),即使修正顺序后,这类值会被分到-1桶,建议补充正负无穷以覆盖所有可能的数值范围。
解决方案
- 将splits调整为严格递增的顺序
- 添加
-float("inf")和float("inf")覆盖全量数值范围
修正后的代码
import pyspark.sql.functions as F from pyspark.ml.feature import Bucketizer df = spark.createDataFrame( [{'score': 0.056906660760916945}, {'score': 0.014312104993006614}, {'score': 0.019505725666714737}, {'score': 0.05695818453036461}, {'score': 0.004712143467991172}, {'score': 0.01950581558579922}, {'score': 0.0}, {'score': 0.004479459183469148}, {'score': 0.0}, {'score': 0.0002634215537602458}, {'score': 1.6}, {'score': 0.0002634215537602458}], ["score"] ) # 调整为严格递增序列,同时覆盖所有数值范围 splits = [-float("inf"), 0.21, 0.7, 1.0, float("inf")] bucketizer = Bucketizer(splits=splits, inputCol="score", outputCol="score_bucket") bucket_df = bucketizer.transform(df.where(F.col("score").isNotNull())) bucket_df.groupBy("score_bucket").count().show()
运行结果
+-------------+-----+ |score_bucket|count| +-------------+-----+ | 0.0| 10| | 3.0| 1| +-------------+-----+
各桶对应区间说明:
- 桶0:
(-inf, 0.21],包含所有小于等于0.21的score值 - 桶1:
(0.21, 0.7],当前数据无此区间的记录 - 桶2:
(0.7, 1.0],当前数据无此区间的记录 - 桶3:
(1.0, inf],包含score=1.6的记录
内容的提问来源于stack exchange,提问作者HeyWatchThis
相关产品推荐
相关产品推荐

