如何子类化tf.train.Saver()?解决Colab训练checkpoint云同步问题
解决Colab训练中Checkpoint自动同步到GCS的问题
Looks like you hit a type check issue in slim.learning.train() — that function expects the saver parameter to be an instance of tf.train.Saver, and your custom callback saver isn't passing that check. Let's break down how to fix this, with two solid approaches:
方案一:继承tf.train.Saver并重写save方法
直接继承原生的Saver类,在保留原有保存逻辑的基础上,添加复制到GCS的步骤。这样就能通过slim的类型检查,同时实现你的需求。
import tensorflow as tf from tensorflow.python.training import saver as tf_saver class GCSSaver(tf_saver.Saver): def __init__(self, gcs_save_dir, *args, **kwargs): super().__init__(*args, **kwargs) # 你的GCS存储目录,比如 'gs://your-bucket-name/checkpoints/' self.gcs_save_dir = gcs_save_dir # 确保GCS目录存在 if not tf.io.gfile.exists(gcs_save_dir): tf.io.gfile.makedirs(gcs_save_dir) def save(self, sess, save_path, global_step=None, latest_filename=None, meta_graph_suffix='meta', write_meta_graph=True, write_state=True): # 先调用原生方法保存到Colab本地 local_save_path = super().save( sess, save_path, global_step, latest_filename, meta_graph_suffix, write_meta_graph, write_state ) # 复制所有checkpoint文件到GCS # 匹配基础文件名的所有关联文件 base_filename = save_path.split('/')[-1] local_checkpoint_files = tf.io.gfile.glob(f"{save_path}*") for local_file in local_checkpoint_files: # 提取文件名,拼接GCS路径 filename = local_file.split('/')[-1] gcs_file_path = f"{self.gcs_save_dir}{filename}" # 覆盖已存在的文件 tf.io.gfile.copy(local_file, gcs_file_path, overwrite=True) return local_save_path
使用方式
# 1. 先授权访问GCS from google.colab import auth auth.authenticate_user() # 2. 构建你的模型和训练操作... train_op = ... # 3. 初始化自定义Saver gcs_saver = GCSSaver( gcs_save_dir='gs://your-bucket/checkpoints/', var_list=tf.trainable_variables() ) # 4. 启动训练 slim.learning.train( train_op, logdir='/tmp/local_checkpoints', # 本地临时存储目录 saver=gcs_saver, save_interval_secs=300, # 每5分钟保存一次 save_summaries_secs=60 )
方案二:使用SessionRunHook(推荐)
如果不想修改Saver的继承关系,可以用TensorFlow标准的SessionRunHook来实现定时/定步保存并同步到GCS。这种方式更符合TensorFlow的扩展设计,也不会破坏原生Saver的逻辑。
import tensorflow as tf class GCSSaveHook(tf.train.SessionRunHook): def __init__(self, saver, local_save_path, gcs_save_dir, save_steps=None, save_secs=None): self.saver = saver self.local_save_path = local_save_path # 本地保存路径,比如 '/tmp/local_checkpoints/model' self.gcs_save_dir = gcs_save_dir self.save_steps = save_steps # 每多少步保存一次 self.save_secs = save_secs # 每多少秒保存一次 self.last_save_time = None self.last_save_step = None # 确保至少设置一种触发条件 if self.save_steps is None and self.save_secs is None: raise ValueError("必须设置save_steps或save_secs中的一个") def begin(self): # 获取全局步数张量 self.global_step_tensor = tf.train.get_or_create_global_step() # 确保GCS目录存在 if not tf.io.gfile.exists(self.gcs_save_dir): tf.io.gfile.makedirs(self.gcs_save_dir) def before_run(self, run_context): # 每次运行前获取当前全局步数 return tf.train.SessionRunArgs(self.global_step_tensor) def after_run(self, run_context, run_values): current_step = run_values.results current_time = tf.timestamp().numpy() # 检查是否满足保存条件 need_save = False if self.save_steps is not None: if self.last_save_step is None or (current_step - self.last_save_step) >= self.save_steps: need_save = True if self.save_secs is not None: if self.last_save_time is None or (current_time - self.last_save_time) >= self.save_secs: need_save = True if need_save: # 保存到本地 self.saver.save( run_context.session, self.local_save_path, global_step=current_step ) # 同步到GCS local_checkpoint_files = tf.io.gfile.glob(f"{self.local_save_path}*") for local_file in local_checkpoint_files: filename = local_file.split('/')[-1] gcs_file_path = f"{self.gcs_save_dir}{filename}" tf.io.gfile.copy(local_file, gcs_file_path, overwrite=True) # 更新最后保存的时间和步数 self.last_save_time = current_time self.last_save_step = current_step
使用方式
# 1. 授权GCS访问 from google.colab import auth auth.authenticate_user() # 2. 构建模型和训练操作... train_op = ... # 3. 初始化原生Saver和自定义Hook saver = tf.train.Saver(var_list=tf.trainable_variables()) gcs_save_hook = GCSSaveHook( saver=saver, local_save_path='/tmp/local_checkpoints/model', gcs_save_dir='gs://your-bucket/checkpoints/', save_secs=300 # 每5分钟保存一次 ) # 4. 启动训练,将Hook传入hooks参数 slim.learning.train( train_op, logdir='/tmp/local_checkpoints', saver=saver, hooks=[gcs_save_hook], save_summaries_secs=60 )
额外提示
- 推荐使用
tf.io.gfile而不是subprocess调用gsutil,因为它是TensorFlow原生的文件操作API,更稳定且不需要额外依赖。 - 如果你的训练代码可以迁移到Keras风格,直接使用
tf.keras.callbacks.ModelCheckpoint并指定GCS路径即可,TensorFlow会自动将Checkpoint保存到GCS,无需手动复制:checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( 'gs://your-bucket/checkpoints/model-{epoch:02d}', save_freq='epoch' ) model.fit(..., callbacks=[checkpoint_callback])
内容的提问来源于stack exchange,提问作者michael
相关产品推荐
相关产品推荐

