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
故障根因
- 核心问题:BatchNormalization层的非训练权重未参与联邦聚合
MobileNetV2包含大量BatchNormalization(BN)层,这类层除了gamma、beta两类可训练参数外,还有两类不可训练的缓冲区权重:滑动平均均值、滑动平均方差,是推理阶段做特征归一化的核心参数。- 训练阶段:TFF默认的联邦平均算法仅聚合客户端上传的可训练权重,BN层的滑动统计量不会参与跨端聚合,不会更新到全局state中;客户端本地训练时,BN层直接使用当前批次的实时统计量计算前向传播,因此训练集上的准确率、损失指标表现正常。
- 评估阶段:
state.model.assign_weights_to()仅会将全局state中存储的可训练权重赋值给中心化Keras模型,BN层的滑动均值、滑动方差仍为模型初始化时的随机值。推理时BN层直接使用这些错误的统计量做归一化,导致模型输出完全失真,最终表现为10分类任务上的随机猜测结果。
- 次要问题:输入预处理不符合模型预期
MobileNetV2官方实现要求输入像素值归一化到[-1, 1]区间,当前代码直接输入0-255范围的原始CIFAR-10像素值,会拖慢模型收敛速度,也会加剧训练、推理阶段的分布偏差。
修复方案
- 针对BN层统计量问题:
- 快速验证方案:评估前遍历模型所有BN层,强制推理时使用当前批次的统计量而非存储的滑动值,即可快速验证权重本身是有效的;
- 规范修复方案:自定义联邦平均聚合逻辑,将BN层的滑动均值、滑动方差纳入待聚合的权重列表,每轮全局更新时同步聚合所有客户端本地更新的BN统计量;或在每轮全局聚合后,用少量校准数据跑一次前向传播,重新计算BN层的滑动统计量。
- 针对预处理问题:在训练、测试的数据流中统一增加预处理逻辑,将像素值从0-255归一化到
[-1, 1]区间,和MobileNetV2的输入要求对齐。
内容的提问来源于stack exchange,提问作者Christian Fachola
相关产品推荐
相关产品推荐

