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

TensorFlow Federated联邦重建模型及权重提取保存问题咨询

解决TensorFlow Federated联邦重建模型权重保存问题

TFF中的ModelWeights是框架自定义的权重结构,并非原生TensorFlow/Keras的模型或权重对象,所以直接调用save()或save_weights()会触发属性错误。要将其用于普通DNN/Sequential模型,需要先把TFF权重转换为原生TF可识别的格式,再赋值给对应结构的Keras模型,最后保存。

具体步骤:

  1. 提取TFF模型权重的原生数值
    从训练后的server_state(教程中联邦训练结束后的服务端状态)中取出权重,将每个层的权重转换为numpy数组:

    # 假设server_state是联邦训练完成后的服务端状态
    tff_weights = server_state.model_weights
    
    # 提取用户嵌入和物品嵌入的权重(对应教程中的矩阵分解模型结构)
    user_emb_weights = tff_weights.user_embedding.embeddings.numpy()
    item_emb_weights = tff_weights.item_embedding.embeddings.numpy()
    
  2. 构建与TFF模型结构一致的Keras模型
    按照教程中矩阵分解模型的结构,创建原生Keras模型:

    import tensorflow as tf
    
    def build_keras_recommender(user_count, item_count, embedding_dim):
        user_input = tf.keras.layers.Input(shape=(1,), name='user_input')
        item_input = tf.keras.layers.Input(shape=(1,), name='item_input')
        
        # 嵌入层要和TFF模型的维度完全匹配
        user_emb = tf.keras.layers.Embedding(
            user_count, embedding_dim, name='user_embedding'
        )(user_input)
        item_emb = tf.keras.layers.Embedding(
            item_count, embedding_dim, name='item_embedding'
        )(item_input)
        
        # 计算点积作为推荐得分
        dot_product = tf.keras.layers.Dot(axes=2)([user_emb, item_emb])
        output = tf.keras.layers.Flatten()(dot_product)
        
        return tf.keras.Model(inputs=[user_input, item_input], outputs=output)
    
    # 替换为教程中实际的用户数、物品数、嵌入维度
    keras_model = build_keras_recommender(
        user_count=1000,
        item_count=500,
        embedding_dim=64
    )
    
  3. 赋值权重并保存
    将提取的权重赋值给Keras模型的对应层,然后保存模型或权重:

    # 给嵌入层设置权重
    keras_model.get_layer('user_embedding').set_weights([user_emb_weights])
    keras_model.get_layer('item_embedding').set_weights([item_emb_weights])
    
    # 保存完整模型
    keras_model.save('federated_recommender_model.h5')
    
    # 或仅保存权重
    keras_model.save_weights('federated_recommender_weights.h5')
    

关键说明:

  • 必须保证Keras模型的结构(层数、维度、层类型)和TFF教程中的模型完全一致,否则权重无法匹配。
  • 若模型包含更多层(如DNN的全连接层),只需按相同方式提取对应层的权重(如tff_weights.dense_layer.kernel.numpy())并赋值给Keras模型的对应层即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 12:35:28