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

TensorFlow Federated Learning训练准确率0.5无提升如何解决

TensorFlow Federated 准确率长期维持0.5问题排查方案

运行配置为10本地epoch、100联邦轮次的TFF训练任务时,模型准确率无提升、始终停留在0.5左右,可按以下优先级定位修复:

1. 模型结构存在硬伤

模型结构张量连接、激活配置错误是首要诱因

  • 现有代码的create_compiled_keras_model函数存在张量断裂问题:第一层GlobalAveragePooling2D的输出未传入下一层,第二层Dense(256)直接跳过池化层接入外部未定义的output张量,池化计算结果完全被丢弃
  • 函数内未显式定义模型输入,直接引用外部作用域的model.input、output变量,极易出现维度不匹配、计算图不连通问题
  • 二分类输出层错误使用relu激活:relu输出范围为[0, +∞),无法输出交叉熵损失要求的概率分布,直接导致损失计算失效
  • 隐藏层Dense(256)未配置激活函数,仅做线性变换,模型拟合能力严重不足

修正后的模型构建代码参考:

def create_compiled_keras_model():
    # 替换为实际任务的输入维度
    model_input = tf.keras.Input(shape=(img_height, img_width, channel_num))
    # 按顺序连接层,禁止跳层引用外部张量
    x = tf.keras.layers.GlobalAveragePooling2D()(model_input)
    x = tf.keras.layers.Dense(units=256, activation='relu')(x)
    # 若使用CategoricalCrossentropy,输出层用softmax激活,输出2维类别概率
    model_output = tf.keras.layers.Dense(units=2, activation='softmax')(x)
    model = tf.keras.Model(model_input, model_output)
    return model

2. 训练与评估配置不匹配

训练、评估阶段的损失、指标、模型结构不一致,会直接导致结果完全失真

  • 训练阶段model_fn使用CategoricalCrossentropy搭配CategoricalAccuracy,评估阶段却切换为BinaryCrossentropy和普通Accuracy:前者要求one-hot标签、概率输出,后者适用于0/1整数标签、类别预测值,配置完全不兼容
  • 普通tf.keras.metrics.Accuracy()不会自动对模型输出的概率值做argmax转换,直接计算会得到完全错误的准确率结果
  • 训练阶段传入TFF的Keras模型未提前编译,未显式指定优化器,容易出现优化器状态初始化异常
  • 评估阶段调用create_keras_model()构建待评估模型,和训练阶段使用的create_compiled_keras_model()结构不一致,assign_weights_to无法正确加载训练权重,评估对象本质是随机初始化模型,准确率自然停留在随机猜测的0.5水平

修正后的模型封装与评估代码参考:

def model_fn():
    keras_model = create_compiled_keras_model()
    # 训练阶段提前编译模型,配置和评估阶段完全对齐
    keras_model.compile(
        optimizer=tf.keras.optimizers.Adam(learning_rate=client_lr),
        loss=tf.keras.losses.CategoricalCrossentropy(),
        metrics=[tf.keras.metrics.CategoricalAccuracy()]
    )
    return tff.learning.from_keras_model(
        keras_model,
        input_spec=sample_batch.element_spec, # 传入输入规范而非直接传batch
        loss=tf.keras.losses.CategoricalCrossentropy(),
        metrics=[tf.keras.metrics.CategoricalAccuracy()]
    )

# 评估阶段配置对齐
eval_model = create_compiled_keras_model() # 调用和训练一致的模型构建函数
eval_model.compile(
    optimizer=tf.keras.optimizers.Adam(learning_rate=client_lr),
    loss=tf.keras.losses.CategoricalCrossentropy(), # 和训练损失保持一致
    metrics=[tf.keras.metrics.CategoricalAccuracy()] # 和训练指标保持一致
)

3. 数据与训练流程校验

  • 校验联邦训练数据格式:确认所有客户端训练数据做了归一化(图像像素值从0-255缩放至0-1区间),标签格式和损失函数要求匹配(CategoricalCrossentropy对应one-hot标签,SparseCategoricalCrossentropy对应整数类别标签),排除客户端数据标签全为同一类的极端非独立同分布情况
  • 校验学习率配置:客户端学习率过大会导致每轮本地训练直接冲垮全局模型权重,可先将client_lr调整到1e-3~1e-4区间测试
  • 校验联邦迭代流程:确认构建iterative_process时传入正确的model_fn,服务端优化器学习率未被错误设置为0,每轮评估必须在权重赋值完成后再执行evaluate操作

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 09:39:42