评分后无法展示Spark DataFrame:列x输入类型不兼容
解决Databricks中XGBoost模型评分时int64转int32的类型不兼容问题
问题根源
训练时模型绑定的输入特征列类型为int32(对应Spark的integer类型),但生产数据中对应列被自动推断为int64(Spark的bigint类型)——即使预处理代码一致,数据源的类型差异或Spark自动推断规则的不同,都会导致类型不匹配,触发转换错误。
解决方案
1. 先定位类型差异
执行代码查看生产数据的列类型,对比训练数据集的schema,找出问题列:
spark_df = spark.createDataFrame(table) spark_df.printSchema()
2. 转换生产数据列类型至模型期望的int32
将问题列强制转换为integer类型(对应int32),示例如下:
import mlflow from pyspark.sql.functions import struct, col model_uri = f"models:/{model_name}/1" spark_df = spark.createDataFrame(table) # 转换单个问题列 spark_df = spark_df.withColumn("x", col("x").cast("integer")) # 若多个bigint列需要转换,可批量处理 # for col_name, dtype in spark_df.dtypes: # if dtype == "bigint": # spark_df = spark_df.withColumn(col_name, col(col_name).cast("integer")) predict = mlflow.pyfunc.spark_udf(spark, model_uri, result_type="double") output_df = spark_df.withColumn("prediction", predict(struct(*spark_df.columns))) output_df.show(20)
3. 创建DataFrame时指定schema(从根源避免类型推断错误)
定义明确的schema,强制指定列类型为IntegerType,避免Spark自动推断为BigIntType:
from pyspark.sql.types import StructType, StructField, IntegerType, DoubleType # 按照训练数据的schema定义生产数据的结构 schema = StructType([ StructField("x", IntegerType(), nullable=True), StructField("feature_1", DoubleType(), nullable=True), # 依次定义其他特征列 ]) spark_df = spark.createDataFrame(table, schema=schema)
4. 验证模型的输入签名
可以查看MLflow模型的输入schema,确认模型期望的输入类型:
import mlflow model_uri = f"models:/{model_name}/1" model = mlflow.pyfunc.load_model(model_uri) print(model.metadata.get_input_schema())
内容的提问来源于stack exchange,提问作者Ullsokk
相关产品推荐
相关产品推荐

