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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 15:42:33