TFF的simple_fedavg示例使用CIFAR-100与VGG19时准确率仅1%
问题原因
- 数据集预处理错误
CIFAR100原生图像格式为(32,32,3)的0-255整数类型,你在element_fn中额外加了tf.expand_dims(element['image'], -1)操作,会导致输入形状变为(32,32,3,1),和VGG19要求的(32,32,3)输入形状完全不匹配,同时没有做像素值归一化,深层神经网络无法正常拟合。 - FedAvg流程参数传参错误
你调用simple_fedavg_tff.build_federated_averaging_process时传入了4个重复的tff_model_fn,该方法在v0.19.0版本的入参仅需要1个模型构造函数,后续跟着服务端、客户端优化器函数,多传的参数会导致优化器逻辑完全错位,模型权重根本没有正确更新。 - 超参数配置完全不符合当前任务
- 服务端SGD学习率设置为1过高,会导致全局参数更新时直接偏离最优区间
- 每轮仅采样4个客户端,CIFAR100联邦划分的非独立同分布程度很高,少量客户端的更新方差极大,模型无法收敛
- 总训练轮数仅50轮,哪怕是集中式训练VGG19在CIFAR100上都需要数百轮才能收敛,联邦训练轮数远远不足。
解决方法
- 修复数据集预处理逻辑,将
element_fn修改为以下内容:
def element_fn(element): return collections.OrderedDict( x=tf.cast(element['image'], tf.float32) / 255.0, y=element['label'] )
- 修正FedAvg流程的传参,删除多余的3个
tff_model_fn,参考官方示例的入参格式修改:
iterative_process = simple_fedavg_tff.build_federated_averaging_process( tff_model_fn, server_optimizer_fn, client_optimizer_fn )
- 调整超参数配置:
- 服务端学习率下调到0.1~0.2区间
- 每轮采样客户端数量提升到20~30,降低更新方差
- 总训练轮数调整到至少200轮
- 客户端优化器先替换为SGD,学习率调整到0.01~0.05区间,联邦场景下SGD比Adam稳定性更高
- 先做集中式验证:单独用你的VGG19模型在集中式CIFAR100数据集上训练,确认能达到30%以上的准确率,排除模型本身的结构问题后再接入联邦训练流程。
- 可在客户端训练逻辑中加入梯度裁剪,阈值设为1.0~5.0,避免深层模型训练时出现梯度爆炸问题。
内容的提问来源于stack exchange,提问作者Fanto
相关产品推荐
相关产品推荐

