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

TensorFlow Federated模型评估时准确率卡在0.1问题求助

TensorFlow Federated 训练指标正常但测试准确率固定为0.1问题

问题现象

在TensorFlow Federated(TFF)框架下实现CIFAR-10数据集分类任务:

  • 训练数据通过tff.simulation.datasets.TestClientData构建,模型采用tf.keras.applications.mobilenet_v2.MobileNetV2结构
  • 训练流程已完成模型与数据格式适配、训练状态初始化,每轮调用state.next执行联邦平均训练
  • 训练集指标表现正常:稀疏分类准确率持续上升、损失稳步下降,但测试集指标完全无变化:准确率固定为0.1、损失稳定在2.302左右(即10分类随机猜测的基准值)

评估逻辑为:每轮训练后构建结构、配置完全一致的独立Keras模型,通过state.model.assign_weights_to(keras_model)加载联邦训练得到的权重。已确认联邦模型与评估模型使用完全相同的SparseCategoricalCrossentropy损失、Adam优化器、SparseCategoricalAccuracy指标,且通过打印张量值验证权重赋值操作已生效,未定位到故障点。

训练日志

-------------------------------ROUND 0 ------------------------------------
Initial weights in state: [-0.00721832  0.01982944  0.08157757]

 Train metrics OrderedDict([('sparse_categorical_accuracy', 0.21346854), ('loss', 2.220433), ('num_examples', 49998), ('num_batches', 1563)]), round time 81.32 seconds

After training weights in state: [-0.00548685  0.01842782  0.0898697 ]

Initial weights in dummy model: tf.Tensor([0.12715222 0.00962208 0.1222005 ], shape=(3,), dtype=float32)

Weights in dummy model after assign weights: tf.Tensor([-0.00548685  0.01842782  0.0898697   0.0580997   0.00205497], shape=(5,), dtype=float32)

 Test metrics [2.3025963306427, 0.10000000149011612]
-------------------------------ROUND 1 ------------------------------------
Initial weights in state: [-0.00548685  0.01842782  0.0898697 ]

 Train metrics OrderedDict([('sparse_categorical_accuracy', 0.27069083), ('loss', 1.9941559), ('num_examples', 49998), ('num_batches', 1563)]), round time 80.71 seconds

After training weights in state: [0.00415337 0.01833635 0.1140746 ]

Initial weights in dummy model: tf.Tensor([ 0.03019264 -0.0810149  -0.01419063], shape=(3,), dtype=float32)

Weights in dummy model after assign weights: tf.Tensor([ 0.00415337  0.01833635  0.1140746   0.05260698 -0.00449031], shape=(5,), dtype=float32)

 Test metrics [2.3026626110076904, 0.10000000149011612]
-------------------------------ROUND 2 ------------------------------------
Initial weights in state: [0.00415337 0.01833635 0.1140746 ]

 Train metrics OrderedDict([('sparse_categorical_accuracy', 0.299912), ('loss', 1.8942232), ('num_examples', 49998), ('num_batches', 1563)]), round time 82.39 seconds

After training weights in state: [0.01262705 0.03320389 0.0952585 ]

Initial weights in dummy model: tf.Tensor([0.12944148 0.07921356 0.11308451], shape=(3,), dtype=float32)

Weights in dummy model after assign weights: tf.Tensor([0.01262705 0.03320389 0.0952585  0.02912883 0.00987563], shape=(5,), dtype=float32)

 Test metrics [2.3031070232391357, 0.10000000149011612]

核心实现代码

基础配置代码

EPOCHS = 1
BATCH_SIZE = 32

# ROUND_CLIENTS <= NUM_CLIENTS
ROUND_CLIENTS = 3
NUM_CLIENTS = 3

NUM_ROUNDS = 3

    
def make_client(num_clients,X, y):
    total_image_count = len(X)
    image_per_set = int(np.floor(total_image_count/num_clients))

    client_train_dataset = collections.OrderedDict()
    for i in range(1, num_clients+1):
        client_name = i-1
        start = image_per_set * (i-1)
        end = image_per_set * i

        print(f"Adding data from {start} to {end} for client : {client_name}")
        data = collections.OrderedDict((('label', y[start:end]), ('pixels', X[start:end])))
        client_train_dataset[client_name] = data
    
    train_dataset = tff.simulation.datasets.TestClientData(client_train_dataset)
    
    return train_dataset

(X_train, y_train), (X_test, y_test) = cifar10.load_data()
cifarFedTrain = make_client(NUM_CLIENTS,X_train,y_train)

def map_fn(example):
    return collections.OrderedDict(
      x=example['pixels'], 
        y=example['label']
    )


def client_data(client_id):
    ds = cifarFedTrain.create_tf_dataset_for_client(cifarFedTrain.client_ids[client_id])
    return ds.repeat(EPOCHS).shuffle(500).batch(BATCH_SIZE).map(map_fn)


train_data = [client_data(n) for n in range(ROUND_CLIENTS)]
element_spec = train_data[0].element_spec

OPTIMIZER = tf.keras.optimizers.Adam()
LOSS = tf.keras.losses.SparseCategoricalCrossentropy()
METRICS=[tf.keras.metrics.SparseCategoricalAccuracy()]

def model_fn():
    model = tf.keras.applications.MobileNetV2((32, 32, 3), classes=10, weights=None)
    return tff.learning.from_keras_model(
            model,
            input_spec=element_spec,
            loss=tf.keras.losses.SparseCategoricalCrossentropy(), 
            metrics=[tf.keras.metrics.SparseCategoricalAccuracy()]
            )


trainer = tff.learning.build_federated_averaging_process(model_fn, client_optimizer_fn=lambda:tf.keras.optimizers.Adam())

训练评估循环代码

def evaluate(state, num_rounds=NUM_ROUNDS): 
    state = trainer.initialize()
    
    for i in range(num_rounds):
        print(f"-------------------------------ROUND {i} ------------------------------------")
        
        print("Initial weights in state:",state.model.trainable[0][0][0][0][:3])
        t1 = time.time()
        state, metrics = trainer.next(state, train_data)
        t2 = time.time()
        print('\n Train metrics {m}, round time {t:.2f} seconds'.format(
            m=metrics['train'], t=t2 - t1))
        
        print("\nAfter training weights in state:",state.model.trainable[0][0][0][0][:3])
        
        model = tf.keras.applications.MobileNetV2((32, 32, 3), classes=10, weights=None, classifier_activation='softmax')
        
        OPTIMIZER = tf.keras.optimizers.Adam()
        LOSS = tf.keras.losses.SparseCategoricalCrossentropy()
        METRICS=[tf.keras.metrics.SparseCategoricalAccuracy()]

        model.compile(OPTIMIZER, LOSS, METRICS)
        print("\nInitial weights in dummy model:",model.weights[0][0][0][0][:3])

        
        
        state.model.assign_weights_to(model)  # Update model with the latest parameters
        
        print("\nWeights in dummy model after assign weights:",model.weights[0][0][0][0][:5])
        metrics_test = model.evaluate(test_data, test_labels, verbose = False)
          
        print('\n Test metrics {m}'.format(m=metrics_test))
    return state, model

故障根因

  1. 核心问题:BatchNormalization层的非训练权重未参与联邦聚合
    MobileNetV2包含大量BatchNormalization(BN)层,这类层除了gamma、beta两类可训练参数外,还有两类不可训练的缓冲区权重:滑动平均均值、滑动平均方差,是推理阶段做特征归一化的核心参数。
    • 训练阶段:TFF默认的联邦平均算法仅聚合客户端上传的可训练权重,BN层的滑动统计量不会参与跨端聚合,不会更新到全局state中;客户端本地训练时,BN层直接使用当前批次的实时统计量计算前向传播,因此训练集上的准确率、损失指标表现正常。
    • 评估阶段:state.model.assign_weights_to()仅会将全局state中存储的可训练权重赋值给中心化Keras模型,BN层的滑动均值、滑动方差仍为模型初始化时的随机值。推理时BN层直接使用这些错误的统计量做归一化,导致模型输出完全失真,最终表现为10分类任务上的随机猜测结果。
  2. 次要问题:输入预处理不符合模型预期
    MobileNetV2官方实现要求输入像素值归一化到[-1, 1]区间,当前代码直接输入0-255范围的原始CIFAR-10像素值,会拖慢模型收敛速度,也会加剧训练、推理阶段的分布偏差。

修复方案

  • 针对BN层统计量问题:
    • 快速验证方案:评估前遍历模型所有BN层,强制推理时使用当前批次的统计量而非存储的滑动值,即可快速验证权重本身是有效的;
    • 规范修复方案:自定义联邦平均聚合逻辑,将BN层的滑动均值、滑动方差纳入待聚合的权重列表,每轮全局更新时同步聚合所有客户端本地更新的BN统计量;或在每轮全局聚合后,用少量校准数据跑一次前向传播,重新计算BN层的滑动统计量。
  • 针对预处理问题:在训练、测试的数据流中统一增加预处理逻辑,将像素值从0-255归一化到[-1, 1]区间,和MobileNetV2的输入要求对齐。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 07:15:53