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

TensorFlow Federated自定义权重聚合器实现报错求助

解决TensorFlow Federated自定义权重聚合的AggregationPlacementError问题

问题背景

基于TensorFlow Federated实现FedAvg时,希望使用自定义权重聚合客户端更新(而非默认的样本数量加权),但自定义WeightedAggregationFactory后运行报错:

AggregationPlacementError: The "result" attribute of return type of next_fn must be placed at SERVER, but found {<float32[7],float32,float32[1],float32>}@CLIENTS.

错误原因

  1. 联邦语义错误:tff.federated_map是在每个客户端执行传入的计算,但原custom_weighted_aggregate的逻辑是试图处理所有客户端的values和weights,这不符合TFF的联邦执行模型——每个客户端只能访问自身的更新和权重。
  2. 结果位置错误:tff.federated_map返回的结果仍位于客户端,而聚合流程要求最终的聚合结果必须位于服务器端。

修正后的实现代码

@tff.tensorflow.computation
def client_weighted_update(value, weight):
    # 客户端计算:将本地更新乘以本地权重
    return tf.nest.map_structure(lambda v: v * weight, value)

class CustomWeightedAggregator(tff.aggregators.WeightedAggregationFactory):
    def create(self, value_type, weight_type):
        @tff.federated_computation
        def initialize():
            # 聚合流程无需状态,返回空状态即可
            return tff.federated_value((), tff.SERVER)

        @tff.federated_computation(
            initialize.type_signature.result,
            tff.FederatedType(value_type, tff.CLIENTS),
            tff.FederatedType(weight_type, tff.CLIENTS)
        )
        def next(state, client_updates, client_weights):
            # 1. 每个客户端计算加权后的本地更新
            weighted_updates = tff.federated_map(client_weighted_update, (client_updates, client_weights))
            
            # 2. 服务器端汇总所有客户端的加权更新和权重
            total_weighted_update = tff.federated_sum(weighted_updates)
            total_weight = tff.federated_sum(client_weights)
            
            # 3. 服务器端归一化得到最终聚合结果
            aggregated_result = tf.nest.map_structure(lambda v: v / total_weight, total_weighted_update)
            
            return tff.templates.MeasuredProcessOutput(
                state=state,
                result=tff.federated_value(aggregated_result, tff.SERVER),
                measurements=tff.federated_value((), tff.SERVER)
            )

        return tff.templates.AggregationProcess(initialize, next)

    @property
    def is_weighted(self):
        return True

代码说明

  1. 客户端计算:client_weighted_update让每个客户端将自身的模型更新乘以自定义权重,确保每个客户端只处理自身数据。
  2. 服务器端汇总:通过tff.federated_sum将所有客户端的加权更新和权重分别汇总到服务器,符合TFF的联邦聚合语义。
  3. 服务器端归一化:在服务器端对汇总后的加权更新进行归一化,明确将结果放置在服务器端,解决位置错误问题。

使用自定义聚合器

将自定义聚合器传入build_weighted_fed_avg的aggregator参数(注意不是client_weighting),同时修正原代码中model_fn的传入方式:

trainer = tff.learning.algorithms.build_weighted_fed_avg(
    model_fn=model_fn,  # 传入模型构造函数,而非已实例化的模型
    client_optimizer_fn=client_optimizer,
    server_optimizer_fn=server_optimizer,
    aggregator=CustomWeightedAggregator()
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 19:50:15