FedPara:基于低秩 Hadamard 积的通信高效联邦学习——Flower 基线实现与复现指南

发布时间:2026/9/17 6:03:12
FedPara:基于低秩 Hadamard 积的通信高效联邦学习——Flower 基线实现与复现指南 FedPara基于低秩 Hadamard 积的通信高效联邦学习——Flower 基线实现与复现指南【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower本文以 baselines/fedpara/README.md 为骨架结合仓库内源码与配置系统解析 FedParaLow-rank Hadamard Product这一通信高效联邦学习方法的原理、Flower 基线实现细节、实验配置与复现步骤。读完本文你将掌握 FedPara 的低秩参数化思想卷积层与全连接层、rank 的计算方式、全局/本地参数分离的个性化变体 pFedPara并能用 Hydra Flower Simulation 完整复现 CIFAR-10、CIFAR-100 与 MNIST 上的通信成本对比实验。FedPara 方法核心低秩 Hadamard 积参数化联邦学习FL中客户端与服务器之间频繁的模型上传与下载是主要的通信瓶颈。FedPara 论文提出了一种通信高效的参数化方法将每一层的权重参数重新参数化为低秩权重矩阵的 Hadamard 积逐元素乘积。设某层原始权重为 WFedPara 将其表达为W W1 ⊙ W2其中 ⊙ 表示 Hadamard 积。与传统的低秩分解如 W UVᵀ不同两个低秩矩阵的 Hadamard 积并不受低秩约束的限制因此参数量被压缩的同时仍能保留远大于传统低秩方法的表现力容量。论文主张在保持可比性能的前提下通信成本可降低为原始层的3 到 10 倍论文主张非本仓库实测结论这是传统低秩方法无法达到的该方法还可与其他高效 FL 优化器组合使用。在此基础上论文进一步提出个性化变体pFedPara将参数划分为全局参数global与本地参数local全局参数参与跨客户端聚合本地参数留在客户端本地训练从而适配非独立同分布non-IID数据下的个性化需求。论文主张 pFedPara 在参数量少三倍以上的情况下仍优于同类个性化 FL 方法。在 models.py 中LowRankNN全连接低秩合成与LowRank卷积低秩合成分别实现了上述两个低秩子矩阵的合成逻辑class LowRankNN(nn.Module): def __init__(self, input_, output, rank) - None: super().__init__() self.X nn.Parameter(torch.empty(size(input_, rank)), requires_gradTrue) self.Y nn.Parameter(torch.empty(size(output, rank)), requires_gradTrue) init.kaiming_normal_(self.X, modefan_out, nonlinearityrelu) init.kaiming_normal_(self.Y, modefan_out, nonlinearityrelu) def forward(self): out torch.einsum(yr,xr-yx, self.Y, self.X) return out卷积侧则在LowRank.forward中用 einsum 完成三维合成return torch.einsum(xyzw,xo,yi-oizw, self.T, self.X, self.Y)即由低秩张量 Tlow_rank × low_rank × kernel × kernel与两个投影因子 Xlow_rank × out_channels、Ylow_rank × in_channels合成一个完整的卷积核 W再交给标准F.conv2d执行卷积。基线仓库结构与执行流程本基线位于baselines/fedpara/核心代码集中在 fedpara/ 子目录文件职责main.pyHydra 入口解析配置、加载数据、构造客户端与策略、启动 Flower 模拟models.pyVGG16GroupNorm与 FC 模型、低秩 Conv2d/Linear 层实现client.py标准FlowerClient与个性化PFlowerClientdataset.pyCIFAR-10/100、MNIST 加载与 IID/non-IID 切分dataset_preparation.pyiid、noniid、noniid_partition_loader切分实现server.py服务端评估函数与加权平均聚合utils.py随机种子、结果存盘、通信成本绘图、全局/本地参数键划分strategy.py空占位文件直接使用 Flower 内置策略main.py的执行流程如下打印解析后的配置调用seed_everything(cfg.seed)固定随机种子np.random、torch、random以及 cudnn 确定性模式见 utils.py通过load_datasets(configcfg.dataset_config, num_clients, batch_size)构造每个客户端的DataLoader与共享测试集由client.gen_client_fn(...)生成客户端工厂函数对于pfedpara/fedper算法会为每个客户端维护独立的状态文件client_{cid}.pth路径由state_path决定instantiate(cfg.model)实例化模型取其参数作为initial_parameters实例化策略默认使用flwr.server.strategy.FedAvg见各 YAML 中strategy._target_并挂载weighted_average评估指标聚合或服务端gen_evaluate_fn调用fl.simulation.start_simulation(...)以client_resources指定的 CPU/GPU 配额运行模拟将history保存为 pickle并调用utils.plot_metric_from_history绘制准确率-通信成本曲线。低秩层与 rank 的确定方式卷积层低秩化自定义Conv2dmodels.py在_calc_from_ratio中根据压缩比ratio计算低秩low_rank最小可能秩 r取ceil(sqrt(out_channels))与ceil(sqrt(in_channels))的较小值最大可能秩 r3论文只给出Rank_min - Rank_max的约束描述本实现通过解一元二次方程a·x² b·x c 0求根a kernel²b out_channels in_channelsc -num_target_params/2取向下取整以避免参数量超过原层最终low_rank ceil((1 - ratio) * r ratio * r3)。前向时对两个低秩合成矩阵做 Hadamard 积W W1() * W2()可选add_nonlinearTrue时先过tanh再执行标准卷积。_init_weights会依据param_type把nn.Conv2d替换为自定义Conv2dlowrank 模式或保持原层并做 He 初始化standard 模式。全连接层低秩化Linear模块同样由_calc_from_ratio(ratio, input_, output)计算 rank前向为w self.w1() * self.w2() self.w1() out F.linear(x, w, self.bias)其中w1为全局跨客户端聚合部分w2为本地不聚合部分——这正是 pFedPara 个性化机制在模型结构上的体现。FC是论文arxiv 1602.05629中的 2NN 全连接模型784→200→256→10MNIST 实验即使用该结构。模型参数量与压缩比model_size属性返回可训练参数量 / 1e6模型字节数换算为 MB。按 README 给出的表格γratio与模型规模关系如下参数比γCIFAR-10CIFAR-1001.0原始15.25M15.30M0.11.55M-0.4-4.53M对应配置CIFAR-10 使用model.ratio: 0.1CIFAR-100 使用model.ratio: 0.4见 cifar10.yaml 与 cifar100.yaml。实验设置与超参数任务与模型任务图像分类模型VGG16 Group Normalizationnum_groups2MNIST 实验使用全连接 2NN。数据集与切分数据集类别数分区数IID 切分non-IID 切分CIFAR-1010100随机切分Dirichlet 分布α0.5CIFAR-10010050随机切分Dirichlet 分布α0.5数据集在 dataset.py 中加载CIFAR 类数据会先尝试读取./data/cifar{10,100}/train_{num_clients}_{alpha:.2f}.pkl若不存在则自动下载并执行iid/noniid切分后以 pickle 缓存README 特别提示数据生成对模型收敛至关重要MNIST 则按shard_size每个 shard 300 样本执行noniid_partition_loader切分。CIFAR 数据增强含 RandomCrop(32, padding4) 与 RandomHorizontalFlip并采用 ImageNet 风格均值/方差归一化MNIST 采用 (0.1307,)/(0.3081,) 归一化。训练超参数超参数CIFAR-10 IIDCIFAR-10 Non-IIDCIFAR-100 IIDCIFAR-100 Non-IIDMNIST每轮客户端数Kfraction16168810总轮数T200200400400100本地 SGD epochE1051055Batch sizeB6464646410初始学习率η0.10.10.10.10.1–0.01学习率衰减τ0.9920.9920.9920.9920.999正则系数λ11110学习率在 models.py 的train中按lr eta_l * learning_decay ** (epoch - 1)逐轮衰减其中epoch即服务端轮次curr_round优化器为无动量、无权重衰减的 SGD损失为 CrossEntropyLoss。各配置文件中hyperparams.eta_l与learning_decay一一对应上述表格MNIST FedAvg 为 η0.05、decay1pFedPara 为 η0.01、decay0.999。环境搭建README 假设已安装pyenv、poetry且用 pyenv 安装 Python 3.10.6随后执行# Set Python 3.10 pyenv local 3.10.6 # Tell poetry to use python 3.10 poetry env use 3.10.6 # Install the base Poetry environment poetry install # Activate the environment poetry shell依赖声明见 pyproject.tomlPython 要求3.10, 3.12.0flwr1.9.0含 simulation 扩展、hydra-core1.3.2、matplotlib、tqdmPyTorch 2.8.0 torchvision 0.23.0cu126 构建。注意启动模拟时client_resources.num_cpus与num_gpus决定每个客户端可占用的并行资源GPU 显存紧张时可提高client_resources.num_gpus即让更少客户端并行。运行实验在激活环境后进入baselines/fedpara/fedpara所在目录以 Hydra 方式运行默认配置为 cifar10# 以默认参数运行 fedpara python -m fedpara.main # 增加轮数并改变本地 epoch 数 python -m fedpara.main num_rounds2024 num_epochs1 # 选择参数化方案lowrank默认或 standard原始权重 python -m fedpara.main model.param_typestandard # 选择 IID 或 non-IID 切分默认 non-iid python -m fedpara.main dataset_config.partitioniid # 指定参数压缩比 python -m fedpara.main model.ratio0.1 # 切换到 CIFAR-100 配置 python -m fedpara.main --config-name cifar100所有 CLI 覆盖项与 cifar10.yaml 等配置文件中的键一一对应例如num_rounds、num_epochs、model.param_type、model.ratio、dataset_config.partition、dataset_config.alpha。复现论文曲线multirunmodel.param_typelowrank对应论文图中带-FedPara后缀的曲线standard对应-orig原始曲线# non-IID CIFAR-10lowrank 与 original 对比 python -m fedpara.main --multirun model.param_typestandard,lowrank # non-IID CIFAR-100 对比 python -m fedpara.main --config-name cifar100 --multirun model.param_typestandard,lowrank # IID CIFAR-10 对比 python -m fedpara.main --multirun model.param_typestandard,lowrank num_epochs10 dataset_config.partitioniid # IID CIFAR-100 对比 python -m fedpara.main --config-name cifar100 --multirun model.param_typestandard,lowrank num_epochs10 dataset_config.partitioniidMNIST 个性化对比实验non-IID分别对应三个配置# FedAvg 基线 python -m fedpara.main --config-name mnist_fedavg # FedPer 个性化 python -m fedpara.main --config-name mnist_fedper # pFedPara 个性化 python -m fedpara.main --config-name mnist_pfedpara三个配置的区别在于mnist_fedavg使用 standard FC无低秩与个性化mnist_fedper使用 FC state_path 按fc1键划分全局/本地参数mnist_pfedpara使用model.param_typelowrank、ratio0.5 按w1/w2键划分。个性化机制的参数划分在 utils.py 的get_keys_state_dict中实现FedPer 将不含fc1的键视为本地参数、含fc1的键为全局参数pFedPara 将含w1的键视为全局、含w2的键视为本地。PFlowerClientclient.py据此在set_parameters中把全局参数从服务器下发覆盖到本地状态同时保留每个客户端client_{cid}.pth中持久化的本地参数。通信成本度量与预期结果度量方式本基线采用论文中的通信成本定义FL 评估通常以达到目标精度所需轮数衡量通信成本但本工作改为评估总传输比特大小2 ×参与客户端数×模型大小×轮数在 utils.py 的plot_metric_from_history中横轴通信成本按r_cc round * 2 * model_size(MB) * clients_per_round / 1024累计为 GB纵轴为轮次对应的中心化/分布式评估准确率从而绘制Accuracy vs Communication Cost (GB)曲线。这也解释了为何 lowrank 方案在相同参数量/通信预算下曲线更陡峭。复现的曲线以下图表对应论文 Figure 3 中 CIFAR-10/100IID 与 non-IID以及 MNIST 个性化Figure 5(c) 相关的结果CIFAR-100 IIDCIFAR-100 Non-IID注意MNIST 部分仅 FedAvg 复现了论文中的结果pFedPara 与 FedPer 在本仓库实现中曾遇到收敛挑战详见下文复现注意事项。复现注意事项重要README 明确记录了复现过程中发现的三个关键事实读者在复现前务必留意对初始化的强依赖FedPara 的低秩训练对初始化极其敏感。实验发现采用 Fan-in HeKaiming初始化会使模型无法收敛性能接近随机分类器只有 Fan-out 初始化能取得预期结果。推测原因是 Fan-out 初始化在反向传播过程中保持了方差守恒。因此 models.py 中所有LowRankNN/LowRank参数均显式使用init.kaiming_normal_(..., modefan_out, nonlinearityrelu)。rank 的确定方式论文除给出 Rank_min - Rank_max 的关系式外未对如何精确计算 rank 给出明确指引。本实现自行推导了与论文说明及约束一致的方程——求解一元二次方程得到max_rank并利用论文的 proposition 2 确定min_rank最终线性插值得到实际 rank详见上文_calc_from_ratio。未实现 Jacobian correction由于论文中关于 dual update 原理的 Jacobian correction 部分缺乏明确实现说明本基线未包含该修正。小结FedPara 基线以低秩 Hadamard 积为核心在保持模型容量的同时大幅压缩了联邦学习中每轮传输的参数量配合 pFedPara 的全局/本地参数分离可同时应对通信瓶颈与数据异构个性化两大问题。仓库在 Flower Simulation 框架下提供了完整的可复现实现从 VGG16GN 的低秩卷积改造、rank 的解析求解到 Dirichlet/分片式 non-IID 数据切分、通信成本曲线的绘制均可通过 Hydra 命令行直接复现。复现时请务必保持 Fan-out 初始化、严格使用本仓库的 rank 计算逻辑并留意个性化算法pFedPara/FedPer在部分设置下的收敛敏感性。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询