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

如何保存tff.learning.templates.LearningAlgorithmState并加载恢复联邦学习训练

联邦学习训练状态的保存与恢复实现

一、修改训练函数,完成10轮训练后保存状态

使用TFF提供的tff.serialize_structured方法将LearningAlgorithmState序列化为字节流,再写入本地文件持久化。修改后的训练函数如下:

import time
import tensorflow_federated as tff

def train_and_save(num_rounds=10, save_path='./federated_state.bin'):
    # 初始化训练状态
    state = trainer.initialize()
    
    # 执行10轮训练
    for round_idx in range(num_rounds):
        start_time = time.time()
        state, metrics = trainer.next(state, client_data)
        elapsed_time = time.time() - start_time
        print(f'Round {round_idx+1}: metrics {metrics}, round time {elapsed_time:.2f}s')
    
    # 序列化训练状态并保存到文件
    serialized_state = tff.serialize_structured(state)
    with open(save_path, 'wb') as f:
        f.write(serialized_state)
    print(f'Training state saved to {save_path}')

二、加载状态并从第11轮恢复训练

通过读取本地文件中的字节流,用tff.deserialize_structured反序列化为原类型的LearningAlgorithmState,然后继续执行后续轮次的训练:

def resume_training(start_round=10, additional_rounds=10, save_path='./federated_state.bin'):
    # 读取保存的序列化状态
    with open(save_path, 'rb') as f:
        serialized_state = f.read()
    
    # 反序列化为LearningAlgorithmState对象
    state = tff.deserialize_structured(
        tff.learning.templates.LearningAlgorithmState, 
        serialized_state
    )
    
    # 从第11轮开始训练(start_round为已完成的轮数)
    for round_idx in range(start_round, start_round + additional_rounds):
        start_time = time.time()
        state, metrics = trainer.next(state, client_data)
        elapsed_time = time.time() - start_time
        print(f'Round {round_idx+1}: metrics {metrics}, round time {elapsed_time:.2f}s')
    
    # 可选:保存更新后的状态,方便后续继续训练
    updated_serialized_state = tff.serialize_structured(state)
    with open(save_path, 'wb') as f:
        f.write(updated_serialized_state)
    print(f'Updated training state saved to {save_path}')

关键注意事项

  • 确保恢复训练时使用的trainer实例与保存状态时的训练器结构完全一致(包括模型定义、优化器配置等),否则反序列化会失败。
  • 保存路径可根据需求自定义,需保证程序对目标路径有读写权限。
  • tff.serialize_structured和tff.deserialize_structured是TFF官方推荐的结构化对象序列化方案,能完整保存LearningAlgorithmState中的模型参数、优化器状态等所有组件。

内容的提问来源于stack exchange,提问作者Akash A R

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 07:03:19