如何在TensorFlow Federated中实现自定义编码处理客户端上传权重
TFF v0.19.0 自定义编解码配置方案
1 补全自定义编码阶段的必填方法
你当前实现的SparseTernaryCompressionEncodingStage中所有标记pass的方法必须按规则补全,否则TFF编码框架无法正常识别调用,补全示例如下:
import tensorflow as tf from tensorflow_model_optimization.python.core.internal import tensor_encoding as te class SparseTernaryCompressionEncodingStage(te.core.EncodingStageInterface): AVERAGE = 'average' NEGATIVES = 'negatives' POSITIVES = 'positives' NEW_SHAPE = 'new_shape' ORIGINAL_SHAPE = 'original_shape' def name(self): return "sparse_ternary_compression" def compressible_tensors_keys(self): # 返回所有编码后传输的字段名 return [self.AVERAGE, self.NEGATIVES, self.POSITIVES, self.NEW_SHAPE, self.ORIGINAL_SHAPE] def commutes_with_sum(self): # 必须先解码再求和,不能直接聚合编码后的数据,所以返回False return False def decode_needs_input_shape(self): # 编码结果中已经包含原始shape信息,不需要外部传入 return False def get_params(self): # 无额外编解码参数,返回空的编码、解码参数字典 return {}, {} def encode(self, original_tensor, encode_params): original_shape = tf.shape(original_tensor) tensor = tf.reshape(original_tensor, [-1]) sparsification_rate = int(len(tensor) / 100 * 1) new_shape = tensor.get_shape().as_list() if sparsification_rate == 0: sparsification_rate = 1 mask = tf.cast(tf.abs(tensor) >= tf.math.top_k(tf.abs(tensor), sparsification_rate)[0][-1], tf.float32) inv_mask = tf.cast(tf.abs(tensor) < tf.math.top_k(tf.abs(tensor), sparsification_rate)[0][-1], tf.float32) tensor_masked = tf.multiply(tensor, mask) average = tf.reduce_sum(tf.abs(tensor_masked)) / sparsification_rate negatives = tf.where(compressed_tensor < 0) positives = tf.where(compressed_tensor > 0) return { self.AVERAGE: average, self.NEGATIVES: negatives, self.POSITIVES: positives, self.NEW_SHAPE: new_shape, self.ORIGINAL_SHAPE: original_shape } # 原有decode方法逻辑错误,必须从传入的encoded_tensors取字段,不能直接调用类属性,修正后代码如下 def decode(self, encoded_tensors, decode_params, num_summands=None, shape=None): average = encoded_tensors[self.AVERAGE] negatives = encoded_tensors[self.NEGATIVES] positives = encoded_tensors[self.POSITIVES] new_shape = encoded_tensors[self.NEW_SHAPE] original_shape = encoded_tensors[self.ORIGINAL_SHAPE] decompressed_tensor = tf.zeros(new_shape, tf.float32) average_values_negative = tf.fill([tf.shape(negatives)[0], ], -average) average_values_positive = tf.fill([tf.shape(positives)[0], ], average) decompressed_tensor = tf.tensor_scatter_nd_update(decompressed_tensor, negatives, average_values_negative) decompressed_tensor = tf.tensor_scatter_nd_update(decompressed_tensor, positives, average_values_positive) return tf.reshape(decompressed_tensor, original_shape)
2 构造适配TFF聚合逻辑的编码聚合器
使用TFF内置的EncodedSumFactory包装你的自定义编码器,自动完成客户端编码、服务器端解码的流程插入:
import tensorflow_federated as tff # 定义编码器构造函数,适配任意float32类型的权重张量 def build_encoder(value_type): stage = SparseTernaryCompressionEncodingStage() return te.encoders.as_simple_encoder(stage, value_type) # 构造编码聚合工厂,自动在客户端执行编码、服务器聚合前执行解码 aggregator = tff.aggregators.EncodedSumFactory( encoder_fn=build_encoder, sum_quantization_factor=None )
3 替换FedAvg默认聚合器
在构造联邦平均训练流程时,传入你自定义的聚合器即可实现需求:
iterative_process = tff.learning.build_federated_averaging_process( model_fn=your_model_fn, # 替换为你自己的模型构造函数 client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.01), server_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=1.0), # 指定自定义的聚合工厂,替换默认的权重增量聚合逻辑 model_update_aggregation_factory=aggregator )
验证逻辑
配置完成后运行训练流程即可看到效果:
- 客户端计算完本地权重增量
weights_delta后自动触发encode方法,仅传输5个编码字段 - 服务器收到所有客户端的编码数据后,自动触发
decode方法还原完整权重增量,再执行内置的加权平均聚合逻辑,不需要修改原有聚合代码
内容的提问来源于stack exchange,提问作者Fanto
相关产品推荐
相关产品推荐

