Snowflake中已拟合编码器保存后推理失败的技术求助
问题描述
- 使用Snowpark与Snowflake执行数据科学流程时,通过
OneHotEncoder、OrdinalEncoder完成特征工程,训练XGBRegressor模型后,用PUT方法将编码器与模型保存至Snowflake,创建UDF推理时失败 - 报错信息:
"ufunc 'isnan' not supported for the input types, and the inputs could not be safely coerced to any supported types according to the casting rule ''safe''" - 额外问题:输出数组格式无逗号分隔(如
['A' 'B' 'C'])
环境信息
- Snowflake版本:7.44.2
- Snowpark for Python版本:1.9.0
- 本地Python版本:3.11(无法降级)
- 本地依赖包:pandas 1.5.3、joblib 1.2.0、xgboost 1.7.3、scikit-learn 1.2.2、numpy 1.25.2
解决方案
1. 解决isnan报错问题
该报错源于数据类型不匹配或序列化/反序列化时类型丢失,按以下步骤处理:
统一序列化组件
将编码器、模型打包为字典后序列化,避免单个组件序列化的兼容性问题:import joblib # 假设已训练好encoder_onehot、encoder_ordinal、model pipeline_components = { "onehot_encoder": encoder_onehot, "ordinal_encoder": encoder_ordinal, "model": model } joblib.dump(pipeline_components, "model_components.joblib", compress=3)正确上传至Snowflake阶段
上传到内部阶段并覆盖旧文件:session.file.put("model_components.joblib", "@MY_STAGE", overwrite=True)UDF内显式处理数据类型
在UDF中转换输入类型为模型预期格式,避免自动转换错误:import snowflake.snowpark as snowpark import joblib import numpy as np def predict(session: snowpark.Session, input_features): # 加载组件 session.file.get("@MY_STAGE/model_components.joblib", "/tmp/") components = joblib.load("/tmp/model_components.joblib") # 转换输入为numpy数组并指定类型 input_arr = np.array(input_features, dtype=np.object_) # 拆分序数/独热特征(根据实际特征位置调整) ordinal_features = input_arr[:, :2].reshape(-1, 2) onehot_features = input_arr[:, 2:].reshape(-1, len(input_arr[:, 2:])) # 执行编码 ordinal_encoded = components["ordinal_encoder"].transform(ordinal_features) onehot_encoded = components["onehot_encoder"].transform(onehot_features) # 合并特征并预测 combined_features = np.hstack([ordinal_encoded, onehot_encoded]) prediction = components["model"].predict(combined_features) return float(prediction[0]) # 注册UDF时严格指定依赖版本 session.udf.register( func=predict, name="PREDICT_MODEL", input_types=[snowpark.types.ArrayType(snowpark.types.StringType())], return_type=snowpark.types.FloatType(), packages=["scikit-learn==1.2.2", "xgboost==1.7.3", "numpy==1.25.2", "joblib==1.2.0"] )
2. 解决数组无逗号格式问题
将numpy数组转换为Python标准列表,Snowflake会自动处理为带逗号的正确格式:
# 编码器输出转换为列表 encoded_result = components["onehot_encoder"].transform(input_data).tolist()
UDF返回数组时,确保返回Python列表而非numpy数组。
3. 关键注意事项
- 依赖版本严格对齐:UDF指定的包版本必须与本地训练时完全一致,避免序列化/反序列化冲突
- 优先用numpy处理数据:Snowpark UDF中尽量避免使用pandas,减少类型转换问题
- 本地预测试序列化:在本地先测试加载序列化文件,确保组件可正常使用后再上传
- 可选:使用Snowpark Model Registry:若版本支持,用官方模型注册表管理组件,比手动PUT更可靠
内容的提问来源于stack exchange,提问作者Hillygoose
相关产品推荐
相关产品推荐

