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

TensorFlow Federated中如何保存最优服务器权重并复用

可行实现方案

你通过get_model_weights(state)拿到的是TFF封装的服务器端权重对象,本身只存储权重张量数值、不绑定模型结构,只要先实例化和训练时结构完全一致的Keras模型骨架承接权重,就可以导出为HDF5或Checkpoint格式供后续加载复用,两种落地方式如下:

  • 方案1:保存为HDF5单文件格式(适合跨场景复用、脱离TFF环境部署)
    先在训练循环外初始化一次和训练、评估逻辑完全同源的Keras模型,在损失下降触发保存的分支里,把TFF权重赋值给模型后直接调用Keras原生保存接口即可:

    # 训练循环外提前初始化模型骨架,和传入build_federated_evaluation的model_fn保持完全一致
    base_model = model_fn()
    lowest_loss = float('inf')
    
    for round_num in range(total_train_rounds):
        # 原有联邦训练逻辑,执行后拿到当前轮state、对应轮次损失loss[round_num]
        if loss[round_num] < lowest_loss:
            lowest_loss = loss[round_num]
            model_weights = transfer_learning_iterative_process.get_model_weights(state)
            # 将TFF格式权重赋值给Keras模型
            model_weights.assign_weights_to(base_model)
            # 保存为h5格式单文件
            base_model.save('best_model_lowest_loss.h5')
    

    后续加载复用如果要对接TFF的联邦评估接口,把加载后的Keras模型权重转回TFF格式即可:

    import tensorflow as tf
    from tensorflow_federated.python.learning import ModelWeights
    
    # 直接加载h5文件,不需要额外复现模型结构
    loaded_model = tf.keras.models.load_model('best_model_lowest_loss.h5')
    # 转换为TFF可识别的权重格式
    loaded_tff_weights = ModelWeights(
        trainable=loaded_model.trainable_variables,
        non_trainable=loaded_model.non_trainable_variables
    )
    # 直接传入评估函数即可
    eval_metric = federated_eval(loaded_tff_weights, [fed_valid_data])
    
  • 方案2:保存为TensorFlow Checkpoint格式(适合训练过程中存最优权重,读写速度更快)
    如果不需要单文件自包含格式,用Checkpoint存储权重张量的读写开销远低于HDF5,适合训练过程中频繁触发保存的场景:

    import tensorflow as tf
    
    base_model = model_fn()
    ckpt = tf.train.Checkpoint(model=base_model)
    ckpt_manager = tf.train.CheckpointManager(ckpt, directory='./train_ckpt', max_to_keep=3)
    lowest_loss = float('inf')
    
    for round_num in range(total_train_rounds):
        # 原有联邦训练逻辑
        if loss[round_num] < lowest_loss:
            lowest_loss = loss[round_num]
            model_weights = transfer_learning_iterative_process.get_model_weights(state)
            model_weights.assign_weights_to(base_model)
            # 保存最优权重checkpoint
            ckpt_manager.save(checkpoint_number=round_num)
    

    后续加载时需要先实例化结构完全一致的模型骨架,再读取权重:

    base_model = model_fn()
    ckpt = tf.train.Checkpoint(model=base_model)
    # 自动读取目录下最新的checkpoint,也可以指定具体ckpt路径
    ckpt.restore(tf.train.latest_checkpoint('./train_ckpt')).expect_partial()
    # 转TFF权重做评估的逻辑和HDF5方案完全一致
    

注意:两种方案都要求存储、加载时使用的model_fn返回的模型结构完全匹配,包括层顺序、参数形状、激活函数配置,否则会触发张量形状不匹配的报错。如果后续需要脱离TFF环境做本地推理、部署,优先选HDF5格式,文件自包含不需要额外复现模型结构;如果仅在训练流水线内做权重缓存、评估,优先选Checkpoint格式,性能更好。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 15:27:12