如何保存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
相关产品推荐
相关产品推荐

