Colab联邦学习训练RAM崩溃,如何实现权重保存与断点续训?
联邦学习断点续训的代码修改方案
核心逻辑
在每轮通信回合(comm_round)完成全局模型聚合后,保存全局模型权重和当前回合数;崩溃重启后加载最近的断点数据,从对应回合继续训练,避免从头开始。
具体代码修改
1. 初始化持久化保存路径(适配Colab)
Colab临时存储会在会话结束后丢失,优先挂载Google Drive保存断点:
from google.colab import drive drive.mount('/content/drive') import os # 自定义断点保存目录,可按需修改 save_dir = "/content/drive/MyDrive/fl_checkpoints" os.makedirs(save_dir, exist_ok=True) # 权重文件命名绑定回合数 checkpoint_path = os.path.join(save_dir, "fl_global_weights_{}.h5") # 记录最新回合数的文件 latest_round_file = os.path.join(save_dir, "latest_comm_round.txt")
2. 训练前加载断点
在初始化全局模型后、启动训练循环前,添加断点检测逻辑:
# 初始化你的联邦学习全局模型(原有代码) global_model = build_your_fl_model() # 检查并加载断点 start_round = 0 if os.path.exists(latest_round_file): with open(latest_round_file, "r") as f: latest_round = int(f.read().strip()) global_model.load_weights(checkpoint_path.format(latest_round)) start_round = latest_round + 1 print(f"已加载断点,将从第 {start_round} 个通信回合开始训练")
3. 训练循环中添加断点保存
在全局模型聚合完成后(即客户端权重上传、聚合更新全局模型的代码之后)插入保存逻辑:
total_comm_rounds = 100 # 你的总训练回合数 for comm_round in range(start_round, total_comm_rounds): # --- 原有训练逻辑:客户端本地训练、权重上传、全局聚合 --- client_weights = train_clients(global_model, client_dataset) aggregated_weights = aggregate_client_weights(client_weights) global_model.set_weights(aggregated_weights) # --- 新增:保存当前回合断点 --- global_model.save_weights(checkpoint_path.format(comm_round)) with open(latest_round_file, "w") as f: f.write(str(comm_round)) # 可选:清理内存缓解Colab RAM压力 import gc gc.collect()
4. 可选:清理旧断点(节省存储空间)
若无需保留所有回合的断点,可在保存新断点时删除旧文件:
# 放在保存新断点的代码之后 if comm_round > 0: old_checkpoint = checkpoint_path.format(comm_round - 1) if os.path.exists(old_checkpoint): os.remove(old_checkpoint)
崩溃后恢复步骤
- 重新挂载Google Drive(若之前挂载过)
- 运行初始化路径、加载断点的代码
- 启动训练循环,自动从断点回合继续执行
内容的提问来源于stack exchange,提问作者Ariaeimehr
相关产品推荐
相关产品推荐

