You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)

崩溃后恢复步骤

  1. 重新挂载Google Drive(若之前挂载过)
  2. 运行初始化路径、加载断点的代码
  3. 启动训练循环,自动从断点回合继续执行

内容的提问来源于stack exchange,提问作者Ariaeimehr

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.14 09:01:14