如何在Azure Notebook中保存lifetimes GammaGamma模型的.pkl文件
在Azure Notebook中保存GammaGamma模型到DBFS或存储容器
保存到DBFS
Databricks File System(DBFS)是Azure Databricks的默认存储,你可以通过两种方式将模型保存到这里:
方式1:直接使用DBFS本地映射路径
DBFS的根目录会挂载到集群节点的/dbfs路径下,你可以直接用这个路径调用save_model:
# 保存模型到DBFS指定路径 model.save_model('/dbfs/models/gammagamma/model.pkl') # 从DBFS加载模型 from lifetimes import GammaGammaFitter loaded_model = GammaGammaFitter() loaded_model.load_model('/dbfs/models/gammagamma/model.pkl')
方式2:通过dbutils操作文件
如果需要更灵活的文件管理,可以先保存到节点本地临时目录,再复制到DBFS:
# 保存到节点本地临时目录 model.save_model('/tmp/model.pkl') # 复制到DBFS目标路径 dbutils.fs.cp('file:/tmp/model.pkl', 'dbfs:/models/gammagamma/model.pkl') # 加载时可将DBFS文件复制到本地再加载 dbutils.fs.cp('dbfs:/models/gammagamma/model.pkl', 'file:/tmp/loaded_model.pkl') loaded_model = GammaGammaFitter() loaded_model.load_model('/tmp/loaded_model.pkl')
保存到Azure Blob存储容器
如果需要将模型保存到外部Azure Blob存储,有两种常用方式:
方式1:将Blob存储挂载到DBFS
先通过Databricks将Blob存储挂载到DBFS(需配置存储账户密钥或SAS令牌),挂载完成后就可以像操作DBFS一样保存模型:
# 假设挂载路径为/dbfs/mnt/yourblobstorage model.save_model('/dbfs/mnt/yourblobstorage/models/gammagamma/model.pkl') # 加载模型 loaded_model = GammaGammaFitter() loaded_model.load_model('/dbfs/mnt/yourblobstorage/models/gammagamma/model.pkl')
方式2:直接使用Azure Storage Blob SDK上传
如果未挂载存储,可通过SDK直接操作Blob存储:
- 先安装依赖库(如果集群未预装):
%pip install azure-storage-blob
- 保存并上传模型:
from azure.storage.blob import BlobServiceClient # 替换为你的Blob存储连接字符串、容器名和目标文件名 connection_string = dbutils.secrets.get("your_secret_scope", "blob_connection_string") # 推荐用秘钥管理,避免硬编码 container_name = "your-container-name" blob_name = "models/gammagamma/model.pkl" # 保存到本地临时文件 model.save_model('/tmp/model.pkl') # 上传到Blob存储 blob_service_client = BlobServiceClient.from_connection_string(connection_string) blob_client = blob_service_client.get_blob_client(container=container_name, blob=blob_name) with open('/tmp/model.pkl', "rb") as data: blob_client.upload_blob(data, overwrite=True)
- 加载模型:
from azure.storage.blob import BlobServiceClient from lifetimes import GammaGammaFitter connection_string = dbutils.secrets.get("your_secret_scope", "blob_connection_string") container_name = "your-container-name" blob_name = "models/gammagamma/model.pkl" # 下载到本地临时文件 blob_service_client = BlobServiceClient.from_connection_string(connection_string) blob_client = blob_service_client.get_blob_client(container=container_name, blob=blob_name) with open('/tmp/loaded_model.pkl', "wb") as download_file: download_file.write(blob_client.download_blob().readall()) # 加载模型 loaded_model = GammaGammaFitter() loaded_model.load_model('/tmp/loaded_model.pkl')
注意:敏感信息(如存储连接字符串)请使用Databricks秘钥范围管理,不要直接硬编码在代码中。
内容的提问来源于stack exchange,提问作者user3490622
相关产品推荐
相关产品推荐

