TensorFlow Federated自定义权重聚合器实现报错求助
解决TensorFlow Federated自定义权重聚合的AggregationPlacementError问题
问题背景
基于TensorFlow Federated实现FedAvg时,希望使用自定义权重聚合客户端更新(而非默认的样本数量加权),但自定义WeightedAggregationFactory后运行报错:
AggregationPlacementError: The "result" attribute of return type of
next_fnmust be placed at SERVER, but found {<float32[7],float32,float32[1],float32>}@CLIENTS.
错误原因
- 联邦语义错误:
tff.federated_map是在每个客户端执行传入的计算,但原custom_weighted_aggregate的逻辑是试图处理所有客户端的values和weights,这不符合TFF的联邦执行模型——每个客户端只能访问自身的更新和权重。 - 结果位置错误:
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
代码说明
- 客户端计算:
client_weighted_update让每个客户端将自身的模型更新乘以自定义权重,确保每个客户端只处理自身数据。 - 服务器端汇总:通过
tff.federated_sum将所有客户端的加权更新和权重分别汇总到服务器,符合TFF的联邦聚合语义。 - 服务器端归一化:在服务器端对汇总后的加权更新进行归一化,明确将结果放置在服务器端,解决位置错误问题。
使用自定义聚合器
将自定义聚合器传入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
相关产品推荐
相关产品推荐

