如何选择最新创建的Keras检查点并上传至GCS
解决检查点排序与GCS上传问题
嘿,这个问题我熟!我来给你拆解一下,从搞定乱序的检查点列表,到把最新的那个传到GCS,一步一步来:
第一步:获取并筛选最新的检查点文件
glob返回的列表确实没有固定顺序,所以我们需要根据文件修改时间或者文件名中的编号来排序,这里推荐两种实用方法:
方法1:按文件修改时间排序(最可靠)
不管文件名怎么命名,最后修改的文件肯定是最新的检查点,这是最稳妥的方式:
import glob import os # 匹配所有检查点文件 model_path_glob = 'checkpoint.*' checkpoints = glob.glob(model_path_glob) # 按文件最后修改时间倒序排列,第一个元素就是最新的 if checkpoints: checkpoints_sorted = sorted(checkpoints, key=lambda f: os.path.getmtime(f), reverse=True) latest_checkpoint = checkpoints_sorted[0] print(f"找到最新检查点:{latest_checkpoint}") else: print("未找到任何检查点文件,请确认路径是否正确")
方法2:按文件名中的编号排序(适合严格命名规则)
如果你的检查点文件名严格遵循checkpoint.xx-{loss}.h5格式,也可以提取编号来排序:
import glob import re model_path_glob = 'checkpoint.*' checkpoints = glob.glob(model_path_glob) # 提取文件名中的数字编号 def get_checkpoint_num(file_path): # 匹配checkpoint.后面的数字部分 num_match = re.search(r'checkpoint\.(\d+)-', file_path) return int(num_match.group(1)) if num_match else 0 if checkpoints: # 按编号倒序排序 checkpoints_sorted = sorted(checkpoints, key=get_checkpoint_num, reverse=True) latest_checkpoint = checkpoints_sorted[0] else: print("未找到检查点文件")
第二步:上传最新检查点到GCS
接下来需要把筛选出来的文件上传到Google Cloud Storage,首先确保你已经安装了google-cloud-storage库:
pip install google-cloud-storage
然后编写上传函数,记得处理GCS的认证(需要你的服务账号密钥文件):
from google.cloud import storage def upload_checkpoint_to_gcs(local_file, bucket_name, gcs_dest_path): """ 上传本地检查点到GCS :param local_file: 本地检查点文件路径 :param bucket_name: GCS桶的名称 :param gcs_dest_path: GCS上的目标路径(比如"models/latest_checkpoint.h5") """ try: # 初始化GCS客户端(如果未设置环境变量,可指定密钥文件路径) # storage_client = storage.Client.from_service_account_json("path/to/your/service-account-key.json") storage_client = storage.Client() # 获取目标桶 bucket = storage_client.get_bucket(bucket_name) # 创建Blob对象(对应GCS上的文件) blob = bucket.blob(gcs_dest_path) # 上传文件,大文件可开启断点续传 blob.upload_from_filename(local_file, resumable=True) print(f"✅ 成功上传 {local_file} 到 GCS:gs://{bucket_name}/{gcs_dest_path}") except Exception as e: print(f"❌ 上传失败:{str(e)}") # 调用示例 if 'latest_checkpoint' in locals(): # 替换成你的GCS桶名称 your_bucket_name = "your-gcs-bucket" # 可以保留原文件名,或者自定义GCS上的路径 gcs_file_path = f"training-checkpoints/{os.path.basename(latest_checkpoint)}" upload_checkpoint_to_gcs(latest_checkpoint, your_bucket_name, gcs_file_path)
注意事项
- 确保你的环境已经配置了GCS的认证权限,可以通过设置
GOOGLE_APPLICATION_CREDENTIALS环境变量,或者在代码中指定服务账号密钥文件。 - 如果检查点文件体积较大,建议开启
resumable=True参数,支持断点续传,避免网络中断导致上传失败。
内容的提问来源于stack exchange,提问作者GRS
相关产品推荐
相关产品推荐

