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

PySpark中如何将Statsmodels训练的模型保存至GCS Bucket?

解决PySpark Pandas UDF中模型保存到GCS的问题

核心问题分析

  1. 直接用Python内置open()写入GCS路径失败:open()仅支持本地文件系统,无法识别GCS的gs://协议。
  2. 保存到工作节点本地:每个节点的本地存储相互独立,Driver端或其他节点无法跨节点访问这些文件。

最优解决方案

方法1:使用gcsfs直接写入GCS(推荐)

gcsfs是适配GCS的Python文件系统库,支持直接读写gs://路径,适合分布式场景。

步骤:

  1. 确保集群所有工作节点安装gcsfs:

    • 提交Spark作业时添加依赖:--packages gcsfs:2023.6.0(替换为对应版本)
    • 或通过集群初始化脚本预先安装。
  2. 修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 05:06:25