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

基于TFF使用VGG16迁移学习时准确率无提升问题求助

问题描述

我在自定义数据集上开展迁移学习,使用ResNet时能取得不错的准确率,但改用VGG16后准确率始终保持不变,损失值却有波动。已采用tf.keras.applications.vgg16.preprocess_input对图像进行预处理,且集中式迁移学习可正常运行。

相关代码

模型定义与联邦学习初始化

if base_model == "VGG16":
    base_model = tf.keras.applications.vgg16.VGG16(
        include_top=False,
        weights="imagenet",
        input_tensor=tf.keras.layers.Input(shape=(input_shape, input_shape, 3)),
        pooling=None,
    )

base_model.trainable = False

inputs = tf.keras.Input(shape=(input_shape, input_shape, 3))
x = base_model(inputs, training=False)

x = tf.keras.layers.GlobalAveragePooling2D()(x)

outputs = tf.keras.layers.Dense(num_classes, activation="softmax")(x)
model = tf.keras.Model(inputs, outputs)

return model

def create_FL_model():
    """create_FL_model_test _summary_

    Returns:
        tff.learning.Model: _description_
    """
    keras_model = load_model(name, base_model)
    return tff.learning.from_keras_model(
        keras_model,
        input_spec=input_spec.element_spec,
        loss=tf.keras.losses.SparseCategoricalCrossentropy(),
        metrics=[
            tf.keras.metrics.SparseCategoricalAccuracy(),
        ],
    )

# 选择联邦学习算法
if fed_alg == "FedAvg":
    transfer_learning_iterative_process = (
        tff.learning.build_federated_averaging_process(
            create_FL_model,
            client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.02),
            server_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=1.0),
        )
    )

keras_model = load_model(name, base_model)

state_transfer = transfer_learning_iterative_process.initialize()

state = tff.learning.state_with_new_model_weights(
    state_transfer,
    trainable_weights=[v.numpy() for v in keras_model.trainable_weights],
    non_trainable_weights=[v.numpy() for v in keras_model.non_trainable_weights],
)

训练循环代码

for epoch in range(num_epochs):
    client_data_train, client_data_valid = client_data.train_test_client_split(
        client_data, num_test_clients=1, seed=12345
    )
    fed_valid_data = preprocess(
        client_data_valid.create_tf_dataset_for_client(
            client_data_valid.client_ids[0]
        )
    )

    random_clients_ids = random.sample(client_data_train.client_ids, k=2)

    federated_train_data = make_federated_data(
        client_data_train, random_clients_ids
    )
    state, metrics = transfer_learning_iterative_process.next(
        state, federated_train_data
    )
排查建议
  • 核对输入尺寸匹配:VGG16的标准输入尺寸为224x224,确认你的input_shape是否满足要求,过小的输入会导致VGG16的特征提取层无法有效输出有区分度的特征。
  • 验证预处理流程:确保tf.keras.applications.vgg16.preprocess_input是图像输入模型前的最后一步预处理,避免后续操作(如额外归一化)破坏VGG16要求的像素值分布([-127.5, 127.5])。
  • 检查权重初始化与更新:确认state_with_new_model_weights传入的可训练权重仅为顶层Dense层的参数,且联邦训练过程中这些权重确实在更新——可以在训练循环中打印顶层权重的变化情况。
  • 调整优化器参数:当前客户端SGD学习率0.02、服务器SGD学习率1.0可能不适合VGG16的特征适配。建议尝试降低客户端学习率至0.001,或改用Adam优化器,观察是否能推动准确率提升。
  • 确认训练模式设置:确保VGG16在联邦训练中始终处于training=False状态,避免BatchNorm层使用客户端本地统计量导致特征混乱。
  • 检查客户端数据分布:若每个选中的客户端数据集标签极端不平衡,联邦平均后的权重更新可能被抵消,导致准确率停滞。可以尝试增加每轮参与训练的客户端数量(如从2增至5-10),或检查客户端数据的标签分布。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 04:48:21