如何通过ML Engine Python客户端API获取GCS上saved_model.pb的路径?
解决GCS上带时间戳的saved_model.pb路径获取问题
我完全懂你的困扰——自动生成的时间戳目录确实会让获取模型路径变得麻烦,不过咱们可以通过Google Cloud Storage的Python客户端库结合正则表达式来搞定,甚至还能直接取最新的那个模型目录,毕竟训练完的最新模型才是你要部署的对吧?
下面给你两种实用的方法:
方法一:遍历GCS目录,筛选并获取最新的时间戳目录
这是最直接可靠的方式,因为我们直接和GCS交互,精准定位到模型文件。
步骤1:安装依赖库
先确保你装了Google Cloud Storage的Python客户端:
pip install google-cloud-storage
步骤2:编写代码获取路径
import re from google.cloud import storage def get_latest_saved_model_path(bucket_name, base_prefix): # 初始化GCS客户端 client = storage.Client() bucket = client.get_bucket(bucket_name) # 列出base_prefix下的所有子目录(也就是那些带时间戳的目录) blobs = bucket.list_blobs(prefix=base_prefix, delimiter='/') timestamp_dirs = [] for prefix in blobs.prefixes: # 提取目录名(去掉末尾的/,然后取最后一段) dir_name = prefix.strip('/').split('/')[-1] # 用正则匹配时间戳格式(通常是纯数字的字符串) if re.fullmatch(r'\d+', dir_name): timestamp_dirs.append(dir_name) if not timestamp_dirs: raise ValueError("没有找到符合格式的时间戳目录") # 按时间戳倒序排序,取最新的那个 timestamp_dirs.sort(reverse=True) latest_dir = timestamp_dirs[0] # 拼接完整的saved_model.pb路径 return f"gs://{bucket_name}/{base_prefix}{latest_dir}/saved_model.pb" # 调用示例 bucket_name = "your-bucket-name" base_prefix = "outputs/export/serv/" saved_model_path = get_latest_saved_model_path(bucket_name, base_prefix) print(f"找到的最新模型路径:{saved_model_path}")
这段代码的逻辑很清晰:
- 先连接到你的GCS存储桶
- 遍历指定前缀下的所有子目录
- 用正则
r'\d+'筛选出纯数字的时间戳目录(如果你的时间戳格式有变化,比如带小数点,只要调整正则就行) - 把目录按倒序排列,取第一个就是最新训练生成的模型目录
- 最后拼接出完整的saved_model.pb路径
方法二:通过ML Engine客户端获取训练任务元数据(可选)
如果你是通过ML Engine的Python客户端提交的训练任务,也可以尝试从训练任务的详情里提取模型导出路径。不过这个方法依赖于任务元数据是否包含完整路径,稳定性不如直接遍历GCS,你可以作为备选:
from google.cloud import ml_v1 def get_model_path_from_job(project_id, job_name): client = ml_v1.JobServiceClient() job_path = client.job_path(project_id, 'us-central1', job_name) job = client.get_job(job_path) # 尝试从任务的输出信息里提取模型路径 # 不同的预构建估算器可能输出字段不同,这里只是示例 if hasattr(job, 'training_output') and job.training_output: export_paths = job.training_output.get('exported_model_paths', []) if export_paths: # 通常exported_model_paths里的路径就是带时间戳的完整目录 return f"{export_paths[0]}/saved_model.pb" return None # 调用示例 project_id = "your-project-id" job_name = "your-training-job-name" saved_model_path = get_model_path_from_job(project_id, job_name) if saved_model_path: print(f"从任务元数据获取的模型路径:{saved_model_path}") else: print("无法从任务元数据获取模型路径,请使用方法一")
注意事项
- 确保你的运行环境有GCS和ML Engine的访问权限,比如设置
GOOGLE_APPLICATION_CREDENTIALS环境变量指向你的服务账号密钥文件 - 如果你的时间戳格式不是纯数字,记得调整正则表达式,比如如果是
YYYYMMDDHHMMSS格式,正则还是r'\d+'就可以,但如果带其他字符,比如20240520-1430,就改成r'\d+-\d+'之类的
内容的提问来源于stack exchange,提问作者DarshanGowda0
相关产品推荐
相关产品推荐

