PySpark中如何将Statsmodels训练的模型保存至GCS Bucket?
解决PySpark Pandas UDF中模型保存到GCS的问题
核心问题分析
- 直接用Python内置
open()写入GCS路径失败:open()仅支持本地文件系统,无法识别GCS的gs://协议。 - 保存到工作节点本地:每个节点的本地存储相互独立,Driver端或其他节点无法跨节点访问这些文件。
最优解决方案
方法1:使用gcsfs直接写入GCS(推荐)
gcsfs是适配GCS的Python文件系统库,支持直接读写gs://路径,适合分布式场景。
步骤:
确保集群所有工作节点安装
gcsfs:- 提交Spark作业时添加依赖:
--packages gcsfs:2023.6.0(替换为对应版本) - 或通过集群初始化脚本预先安装。
- 提交Spark作业时添加依赖:
修改Pandas UDF中的保存代码:
import gcsfs import joblib # 推荐用joblib替代pickle,更适合序列化机器学习模型 def train_and_save_group_model(df): # 1. 训练模型逻辑(示例) from sklearn.linear_model import LinearRegression X = df.drop(['l2_category', 'target'], axis=1) y = df['target'] model = LinearRegression().fit(X, y) # 2. 获取当前分组的类别标识 l2_category = df['l2_category'].iloc[0] gcs_model_dir = "gs://darkstores-data-eng_stg/multi_sku_test/models/" # 3. 用gcsfs写入GCS fs = gcsfs.GCSFileSystem(project="your-gcp-project-id") # 替换为你的GCP项目ID model_path = f"{gcs_model_dir}{l2_category}.pkl" with fs.open(model_path, 'wb') as f: joblib.dump(model, f) # 4. 返回业务结果DataFrame return df[['l2_category', 'prediction']]
方法2:本地临时文件中转后上传(无额外依赖场景)
如果无法安装gcsfs,可以先将模型保存到工作节点临时文件,再通过gsutil命令上传到GCS。
import tempfile import os import joblib def train_and_save_group_model(df): # 训练模型逻辑 model = ... # 你的训练代码 l2_category = df['l2_category'].iloc[0] gcs_model_dir = "gs://darkstores-data-eng_stg/multi_sku_test/models/" # 1. 保存到本地临时文件 with tempfile.NamedTemporaryFile(mode='wb', delete=False) as tmp_file: joblib.dump(model, tmp_file) tmp_path = tmp_file.name # 2. 用gsutil上传到GCS os.system(f"gsutil cp {tmp_path} {gcs_model_dir}{l2_category}.pkl") # 3. 清理临时文件 os.unlink(tmp_path) return df[['l2_category', 'prediction']]
注意:需确保工作节点的服务账号拥有GCS写入权限,且已安装配置gsutil。
方法3:Driver端统一保存(小分组场景)
如果分组数量较少,可在Pandas UDF中序列化模型并返回,再在Driver端收集后统一写入GCS,避免工作节点依赖。
import pickle import pandas as pd def train_group_model(df): # 训练模型逻辑 model = ... # 你的训练代码 l2_category = df['l2_category'].iloc[0] # 序列化模型为二进制数据 model_bytes = pickle.dumps(model) return pd.DataFrame({ 'l2_category': [l2_category], 'model_binary': [model_bytes], 'prediction': df['prediction'] }) # 在Driver端执行分组训练 result_df = df.groupBy('l2_category').applyInPandas( train_group_model, schema="l2_category string, model_binary binary, prediction double" ) # 收集模型并保存到GCS import gcsfs fs = gcsfs.GCSFileSystem(project="your-gcp-project-id") gcs_model_dir = "gs://darkstores-data-eng_stg/multi_sku_test/models/" for row in result_df.collect(): category = row.l2_category model_data = row.model_binary with fs.open(f"{gcs_model_dir}{category}.pkl", 'wb') as f: f.write(model_data)
注意:分组数量过多时,collect()会占用大量Driver内存,不推荐。
关键注意事项
- 二进制写入:序列化模型必须用
wb模式,不能用文本模式w,否则会导致模型损坏或序列化失败。 - 权限配置:确保Spark集群使用的服务账号拥有
storage.objects.create和storage.objects.list权限,可通过GCP IAM配置。 - 依赖一致性:所有工作节点需安装相同版本的机器学习库(如scikit-learn)和
gcsfs,避免反序列化时出现版本兼容问题。
内容的提问来源于stack exchange,提问作者Jack Daniel
相关产品推荐
相关产品推荐

