在Databricks上循环训练SparkXGBClassifier时遇类型错误及阶段失败
在Databricks的i3.2xlarge实例(61GB内存、8核)上通过for循环训练100余个SparkXGBClassifier模型,代码如下:
import os os.environ["PYSPARK_PIN_THREAD"] = "true" spark.conf.set("spark.yarn.executor.memoryOverhead", 600) spark.conf.set("spark.sql.execution.arrow.pyspark.enabled", "true") for i in models: labeled_train_df = get_train_data(df, i) if labeled_train_df: tuned_model = spark.table("default.tuned_params").select("model_name", "xgboost_params").collect() tuned_model_dict = {i.model_name: i.xgboost_params for i in tuned_model} param_dict = tuned_model_dict.get(i) if param_dict: _n_estimators, _learning_rate, _max_depth, _gamma = param_dict['n_estimators'], param_dict['learning_rate'], param_dict['max_depth'], param_dict['gamma'] else: _n_estimators, _learning_rate, _max_depth, _gamma = 100, 0.30, 6, 0 xgb = SparkXGBClassifier( eval_metric='auc', booster='gbtree', features_col='features', label_col='label', tree_method='hist', n_estimators = _n_estimators, learning_rate = _learning_rate, max_depth = _max_depth, gamma = _gamma ) train_data, test_data = labeled_train_df.randomSplit([0.7, .3]) model = xgb.fit(train_data) model.save(f'{s3}/{i}') del model, train_data, test_data, labeled_train_df
第2次迭代后出现如下错误:
[0;31mPy4JJavaError[0m: An error occurred while calling z:org.apache.spark.api.python.PythonRDD.collectAndServe. : org.apache.spark.SparkException: Job aborted due to stage failure: Could not recover from a failed barrier ResultStage. Most recent failure reason: Stage failed because barrier task ResultTask(96121, 0) finished unsuccessfully. org.apache.spark.api.python.PythonException: 'TypeError: 'str' object cannot be interpreted as an integer'. Full traceback below: Traceback (most recent call last): File "/local_disk0/.ephemeral_nfs/cluster_libraries/python/lib/python3.9/site-packages/xgboost/spark/core.py", line 836, in _train_booster booster = worker_train( File "/local_disk0/.ephemeral_nfs/cluster_libraries/python/lib/python3.9/site-packages/xgboost/core.py", line 620, in inner_f return func(**kwargs) File "/local_disk0/.ephemeral_nfs/cluster_libraries/python/lib/python3.9/site-packages/xgboost/training.py", line 182, in train for i in range(start_iteration, num_boost_round): TypeError: 'str' object cannot be interpreted as an integer at org.apache.spark.api.python.BasePythonRunner$ReaderIterator.handlePythonException(PythonRunner.scala:687) at org.apache.spark.sql.execution.python.PythonArrowOutput$$anon$1.read(PythonArrowOutput.scala:118) at org.apache.spark.api.python.BasePythonRunner$ReaderIterator.hasNext(PythonRunner.scala:640) at org.apache.spark.InterruptibleIterator.hasNext(InterruptibleIterator.scala:37) at scala.collection.Iterator$$anon$11.hasNext(Iterator.scala:491) at scala.collection.Iterator$$anon$10.hasNext(Iterator.scala:460) at scala.collection.Iterator$$anon$10.hasNext(Iterator.scala:460) at org.apache.spark.ContextAwareIterator.hasNext(ContextAwareIterator.scala:39) at org.apache.spark.api.python.SerDeUtil$AutoBatchedPickler.hasNext(SerDeUtil.scala:88) at scala.collection.Iterator.foreach(Iterator.scala:943) at scala.collection.Iterator.foreach$(Iterator.scala:943) at org.apache.spark.api.python.SerDeUtil$AutoBatchedPickler.foreach(SerDeUtil.scala:82) at org.apache.spark.api.python.PythonRDD$.writeIteratorToStream(PythonRDD.scala:464) at org.apache.spark.api.python.PythonRunner$$anon$2.writeIteratorToStream(PythonRunner.scala:864) at org.apache.spark.api.python.BasePythonRunner$WriterThread.$anonfun$run$1(PythonRunner.scala:566) at org.apache.spark.util.Utils$.logUncaughtExceptions(Utils.scala:2359) at org.apache.spark.api.python.BasePythonRunner$WriterThread.run(PythonRunner.scala:357) at org.apache.spark.scheduler.DAGScheduler.failJobAndIndependentStages(DAGScheduler.scala:3377) at org.apache.spark.scheduler.DAGScheduler.$anonfun$abortStage$2(DAGScheduler.scala:3309) at org.apache.spark.scheduler.DAGScheduler.$anonfun$abortStage$2$adapted(DAGScheduler.scala:3300) at scala.collection.mutable.ResizableArray.foreach(ResizableArray.scala:62) at scala.collection.mutable.ResizableArray.foreach$(ResizableArray.scala:55) at scala.collection.mutable.ArrayBuffer.foreach(ArrayBuffer.scala:49) at org.apache.spark.scheduler.DAGScheduler.abortStage(DAGScheduler.scala:3300) at org.apache.spark.scheduler.DAGScheduler.handleTaskCompletion(DAGScheduler.scala:2790) at org.apache.spark.scheduler.DAGSchedulerEventProcessLoop.doOnReceive(DAGScheduler.scala:3583) at org.apache.spark.scheduler.DAGSchedulerEventProcessLoop.onReceive(DAGScheduler.scala:3527) at org.apache.spark.scheduler.DAGSchedulerEventProcessLoop.onReceive(DAGScheduler.scala:3515) at org.apache.spark.util.EventLoop$$anon$1.run(EventLoop.scala:51) at org.apache.spark.scheduler.DAGScheduler.$anonfun$runJob$1(DAGScheduler.scala:1178) at scala.runtime.java8.JFunction0$mcV$sp.apply(JFunction0$mcV$sp.java:23) at com.databricks.spark.util.FrameProfiler$.record(FrameProfiler.scala:80) at org.apache.spark.scheduler.DAGScheduler.runJob(DAGScheduler.scala:1166) at org.apache.spark.SparkContext.runJobInternal(SparkContext.scala:2739) at org.apache.spark.rdd.RDD.$anonfun$collect$1(RDD.scala:1070) at org.apache.spark.rdd.RDDOperationScope$.withScope(RDDOperationScope.scala:165) at org.apache.spark.rdd.RDDOperationScope$.withScope(RDDOperationScope.scala:125) at org.apache.spark.rdd.RDDOperationScope$.withScope(RDDOperationScope.scala:112) at org.apache.spark.rdd.RDD.withScope(RDD.scala:445) at org.apache.spark.rdd.RDD.collect(RDD.scala:1068) at org.apache.spark.api.python.PythonRDD$.collectAndServe(PythonRDD.scala:282) at org.apache.spark.api.python.PythonRDD.collectAndServe(PythonRDD.scala) at sun.reflect.GeneratedMethodAccessor663.invoke(Unknown Source) at sun.reflect.DelegatingMethodAccessorImpl.invoke(DelegatingMethodAccessorImpl.java:43) at java.lang.reflect.Method.invoke(Method.java:498) at py4j.reflection.MethodInvoker.invoke(MethodInvoker.java:244) at py4j.reflection.ReflectionEngine.invoke(ReflectionEngine.java:380) at py4j.Gateway.invoke(Gateway.java:306) at py4j.commands.AbstractCommand.invokeMethod(AbstractCommand.java:132) at py4j.commands.CallCommand.execute(CallCommand.java:79) at py4j.ClientServerConnection.waitForCommands(ClientServerConnection.java:195) at py4j.ClientServerConnection.run(ClientServerConnection.java:115) at java.lang.Thread.run(Thread.java:750)
用户疑问
- 这是内存不足(OOM)问题吗?
- 每次迭代仅处理约300万行训练数据,为何手动删除变量后前两次迭代能成功,后续却报错?
- 该如何解决?
1. 不是OOM问题,是参数类型错误
从报错栈核心信息TypeError: 'str' object cannot be interpreted as an integer可以明确,问题出在XGBoost训练时,num_boost_round(对应n_estimators参数)被传入了字符串类型,无法被解析为整数,导致循环range(start_iteration, num_boost_round)执行失败。
前两次迭代成功是因为对应模型的tuned_params表中参数是整数类型,后续迭代的模型对应的n_estimators参数被存储为字符串,才触发了这个错误。手动删除变量只是释放Python端内存,和这个类型错误完全无关。
2. 错误根源
代码从tuned_params表读取xgboost_params字典后,直接取出参数赋值给变量,但如果表中存储的n_estimators、max_depth等参数是字符串类型(比如写入表时未做类型转换),就会导致传递给SparkXGBClassifier的参数类型不符合要求,最终在XGBoost底层触发类型错误。
3. 解决方案
方案一:强制转换参数类型
在读取参数后,显式将整数类型参数转为int,浮点型参数转为float,确保类型正确:
if param_dict: # 强制转换参数类型,避免字符串传入 _n_estimators = int(param_dict['n_estimators']) _learning_rate = float(param_dict['learning_rate']) _max_depth = int(param_dict['max_depth']) _gamma = float(param_dict['gamma']) else: _n_estimators, _learning_rate, _max_depth, _gamma = 100, 0.30, 6, 0
方案二:修正tuned_params表的数据类型
检查tuned_params表中xgboost_params字段的存储格式,确保n_estimators、max_depth是整数类型而非字符串。如果是通过DataFrame写入的表,写入前要保证参数的原始类型正确,避免被自动转为字符串。
方案三:优化参数读取逻辑
每次循环都调用spark.table(...).collect()会重复读取全表数据,建议提前将所有调优参数加载一次,放在循环外部,避免重复IO:
# 提前加载所有调优参数,放在循环外 tuned_model = spark.table("default.tuned_params").select("model_name", "xgboost_params").collect() tuned_model_dict = {row.model_name: row.xgboost_params for row in tuned_model} for i in models: labeled_train_df = get_train_data(df, i) if labeled_train_df: param_dict = tuned_model_dict.get(i) if param_dict: _n_estimators = int(param_dict['n_estimators']) _learning_rate = float(param_dict['learning_rate']) _max_depth = int(param_dict['max_depth']) _gamma = float(param_dict['gamma']) else: _n_estimators, _learning_rate, _max_depth, _gamma = 100, 0.30, 6, 0 # 后续模型训练代码不变
额外优化:释放Spark资源
虽然本次错误不是内存问题,但循环训练模型时,除了删除Python变量,还可以调用spark.catalog.clearCache()清理Spark缓存,避免累积占用内存:
model.save(f'{s3}/{i}') # 清理Python变量和Spark缓存 del model, train_data, test_data, labeled_train_df spark.catalog.clearCache()
内容的提问来源于stack exchange,提问作者WhimsicalWhale

