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

FLAML AutoML预测结果与概率不匹配问题及实现优化问询

问题解答

1. 预测结果与最高概率不匹配的原因

核心原因是数据拼接时的索引错位:

  • 当你把pyspark.pandas.Series(y_pred、y_pred_prob)转成Pandas Series/DataFrame,再和test_df转成的Pandas DataFrame用pd.concat(axis=1)拼接时,三者的索引可能不一致。
  • Spark DataFrame转Pandas时默认生成连续的0起始索引,但pyspark.pandas.Series转Pandas后的索引可能因为Spark分区、数据处理流程的差异,出现不连续或错位的情况。最终拼接后,某一行的预测值对应的是另一行的概率值,导致看起来预测结果和最高概率不匹配。

2. 更简洁的实现方法

尽量避免频繁在Spark和Pandas之间切换,直接在pyspark.pandas层面完成数据整合,保证索引一致性:

# 保留测试集需要的所有列(ID列、目标列、特征列),转成pyspark.pandas DataFrame
psdf_test = to_pandas_on_spark(test_df.select("pdate", "zm", "alpha", "features"))

# 直接添加预测列
psdf_test["prediction"] = automl.predict(psdf_test[["features"]])

# 处理概率列:可选择拆分为单独类别列或保留为列表列
# 方式1:拆分为对应类别0、1、2的概率列
psdf_test["prob_0"] = psdf_test.apply(lambda row: automl.predict_proba(row[["features"]])[0][0], axis=1)
psdf_test["prob_1"] = psdf_test.apply(lambda row: automl.predict_proba(row[["features"]])[0][1], axis=1)
psdf_test["prob_2"] = psdf_test.apply(lambda row: automl.predict_proba(row[["features"]])[0][2], axis=1)

# 方式2:保留概率列表列
psdf_test["probability"] = automl.predict_proba(psdf_test[["features"]])

# 按需转换格式
result_pd = psdf_test.to_pandas()  # 转Pandas DataFrame
# result_spark = psdf_test.to_spark()  # 转Spark DataFrame

也可以直接用Spark原生API操作,跳过pyspark.pandas转换:

# 用最优lgbm_spark模型直接对Spark DataFrame预测
predictions_spark = automl.model.predict(test_df)
probabilities_spark = automl.model.predict_proba(test_df)

# 合并原测试集、预测结果与概率列
from pyspark.sql.functions import col

result_spark = test_df.join(
    predictions_spark.withColumnRenamed("prediction", "prediction"),
    on=["pdate", "zm"]
).join(
    probabilities_spark.select(
        col("pdate"),
        col("zm"),
        col("probability")[0].alias("prob_0"),
        col("probability")[1].alias("prob_1"),
        col("probability")[2].alias("prob_2")
    ),
    on=["pdate", "zm"]
)

# 转Pandas(如果需要)
result_pd = result_spark.toPandas()

内容的提问来源于stack exchange,提问作者Kelaref

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 11:45:03