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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 18:06:07