Flower Baseline 实践:在 Flower 中复现 DASHA 分布式非凸优化与通信压缩算法
2026/9/16 18:40:15 网站建设 项目流程

Flower Baseline 实践:在 Flower 中复现 DASHA 分布式非凸优化与通信压缩算法

【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower

DASHA(Distributed nonconvex optimization with communication compression and optimal oracle complexity)是由 Alexander Tyurin 与 Peter Richtárik 提出的分布式非凸优化方法家族,在联邦学习场景下通过"只传输压缩向量 + 方差缩减"的机制同时获得最优的 oracle 复杂度与通信复杂度。本文基于 Flower 框架下的 DASHA baseline,完整讲解该 baseline 的实验设置、环境搭建、运行命令与源码实现,读者可以据此在 Flower 中一键复现 DASHA 与 MARINA 的对比实验,并理解压缩通信策略在服务端与客户端两侧的落地细节。

论文背景:DASHA 方法家族的核心思想

DASHA 针对的是分布式非凸优化问题:各节点上的局部目标函数具有有限和(finite-sum)或期望(expectation)形式,节点之间通过通信交换信息来协同最小化全局目标。论文提出了一个新的方法家族,包括DASHA-PAGE、DASHA-MVR 与 DASHA-SYNC-MVR,其核心改进在于:

  • 相比此前 SOTA 方法 MARINA(Gorbunov et al., 2020),DASHA 系列在理论上改进了 oracle 复杂度与通信复杂度;
  • 以随机稀疏化算子 RandK 为例,为达到 ε-平稳点,有限和情形下方法只需计算O(√m / (ε√n))个梯度,期望形式下为O(σ / (ε^{3/2} n))个梯度,同时保持了 SOTA 的通信复杂度O(d / (ε√n))
  • 与 MARINA 不同,DASHA、DASHA-PAGE 与 DASHA-MVR 只发送压缩后的向量,因此对联邦学习场景更实用;
  • 论文还将结果推广到满足 Polyak-Lojasiewicz 条件的函数,并在非凸分类与深度学习模型训练实验中得到显著改进验证。

如果您的项目使用本 baseline,请记得同时引用原论文作者与 Flower 论文。

Baseline 概览:实现了什么

本目录(baselines/dasha)实现了 DASHA 论文中的实验,具体包含:

  • 实现内容:DASHA 论文实验的完整复现,同时提供 MARINA 作为对照基线;
  • 数据集:LIBSVM 的 mushrooms 数据集与 PyTorch Torchvision 的 CIFAR10 数据集;
  • 硬件建议:原实验在一台 64 核桌面机器上运行。任何 1 核机器即可运行 mushrooms 实验;CIFAR10 实验需要更多 CPU 资源(例如 4 核即可满足)以及 1 块支持 CUDA 的 GPU;
  • 贡献者:Alexander Tyurin。

实验设置

任务与模型

Baseline 覆盖两类任务:

  • 图像分类(CIFAR10);
  • 线性回归(mushrooms,使用论文 Section A.1 的非凸损失)。

对应的两个模型实现位于 models.py:

模型说明配置
LinearNetWithNonConvexLoss逻辑回归模型 + 论文 Section A.1 的非凸损失conf/model/linear_net_with_non_convex_loss.yaml
ResNet18WithLogisticLossResNet18 网络 + 交叉熵损失(论文 Section A.4)conf/model/resnet_18_with_logistic_loss.yaml

值得说明的是,论文中使用的非凸损失(NonConvexLoss)在 models.py 中有明确实现:它将目标标签映射到{-1, +1},计算sigmoid(w·x·y)后取(1 - sigmoid)^2的均值,这一设计正是为了在简单模型上验证算法处理非凸目标的收敛行为。

数据集与划分方式

默认数据集按**随机划分(random)**方式切分给n个客户端:

数据集类别数划分方式
mushrooms2random
CIFAR1010random

数据集的加载逻辑在 dataset.py 中:mushrooms 通过sklearn.datasets.load_svmlight_file读取 LIBSVM 格式文件并将标签重映射为 0/1;CIFAR10 使用 Torchvision 加载并做ToTensor + Normalize预处理。切分则使用torch.utils.data.random_split,将整个训练集按1/num_clients等分给各客户端(dataset.py 的random_split)。若未指定path_to_dataset,dataset_preparation.py 会自动下载数据集到默认路径(mushrooms 的下载地址记录在 conf/dataset/libsvm.yaml 的_dataset_urls中)。

训练超参数

所有实验中,算法参数均取自论文理论推荐值,唯一步长(step size)需要调节:

  • mushrooms 实验:步长从{0.25, 0.5, 1.0}(2 的幂集合)中微调;
  • CIFAR10 实验:步长固定为0.01

环境搭建

Baseline 基于 Poetry 管理依赖,Python 版本限定为>=3.10.0, <3.11.0(见 pyproject.toml)。按以下步骤构建环境:

# 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 # By default, Poetry installs the PyTorch package with Python 3.10 and CUDA 11.8. # If you have a different setup, then change the "torch" and "torchvision" lines in [tool.poetry.dependencies]. poetry install # Activate the environment poetry shell

主要依赖包括flwr(含 simulation 扩展)、hydra-corescikit-learnmatplotlibtorch/torchvision。PyTorch 的安装方式因平台而异:Linux 默认安装 CUDA 11.8 版(torch-2.0.0+cu118),macOS 则安装 CPU 版,若环境不同请相应修改 pyproject.toml 中[tool.poetry.dependencies]torchtorchvision行。

运行实验

激活 Poetry 环境(在 baselines/dasha 目录下执行poetry shell)后,即可运行默认配置:

python -m dasha.main # this will run using the default settings in `dasha/conf`

默认配置定义在 conf/base.yaml 中:num_clients: 5num_rounds: 10000、使用RandKCompressor压缩器(number_of_coordinates: 1),默认数据集为 libsvm(mushrooms)、默认方法为 dasha、默认模型为线性非凸损失模型。

命令行覆盖配置

Hydra 支持直接从命令行覆盖任意配置项:

# The following commands runs an experiment with the step size 0.5. # Instead of the full, non-compressed vectors, each node sends a compressed vector with only 10 coordinates. python -m dasha.main method.strategy.step_size=0.5 compressor.number_of_coordinates=10 # if you run this baseline with a larger model, you might want to use the GPU (not used by default). python -m dasha.main method.client.device=cuda

常用配置项说明:

配置路径含义默认值
method.strategy.step_size服务端聚合时的更新步长dasha/marina 为 0.5,stochastic 变体为必填(???
compressor.number_of_coordinatesRandK 压缩器每轮保留的坐标数 K1
method.client.device客户端训练设备cpu
method.client.send_gradient是否让客户端在 evaluate 阶段回传完整梯度(用于计算梯度范数指标)false
method.client.mega_batch_size随机化变体的 mega-batch 大小(计算初始梯度时用)dasha: 100,marina: 10
method.client.batch_size随机化变体采样的小批量大小25
method.client.strict_load加载服务端参数时是否严格要求无缺失键true
dataset=cifar10切换数据集libsvm(mushrooms)
num_rounds联邦训练轮数10000
local_address服务端监听地址(多进程并行运行时使用)localhost:8080

运行 MARINA 对照

论文以 MARINA(Gorbunov et al., 2020)为对照基线,切换方法只需一行:

python -m dasha.main method=marina

method配置组(conf/method)内置了四种方法:

  • dasha:确定性 DASHA(dasha.yaml),客户端为DashaClient
  • marina:确定性 MARINA(marina.yaml),客户端为MarinaClient
  • stochastic_dasha:随机化 DASHA(stochastic_dasha.yaml),客户端为StochasticDashaClient
  • stochastic_marina:随机化 MARINA(stochastic_marina.yaml),客户端为StochasticMarinaClient

其中stochastic_dasha引入了stochastic_momentum: 0.1这一额外的动量超参数,对应论文中处理期望形式目标的方差缩减设计。

源码级解析:压缩通信如何在 Flower 中落地

了解底层实现有助于正确调参。整个 baseline 以"服务端启动、客户端多进程并行"的方式运行:入口 main.py 通过multiprocessing启动num_clients + 1个进程,进程 0 运行 Flower 服务端(fl.server.start_server),其余进程各自运行一个 Flower 客户端(fl.client.start_numpy_client)并连接到local_address

服务端:梯度估计器与参数更新(strategy.py)

服务端逻辑集中在 strategy.py 的_CompressionAggregator中,其注释明确指出该实现对应DASHA 论文 Algorithm 1(MARINA 的逻辑几乎相同):

  • 服务端维护全局参数_parameters与梯度估计器_gradient_estimator
  • 每轮收集各客户端返回的压缩向量,先估算每个客户端收到的比特数(estimate_size),再解压并取平均;
  • _gradient_estimator为 None(首轮),则直接将该均值作为初始梯度估计器;否则累加;
  • 最后执行_parameters -= step_size * gradient_estimator完成一步更新。

DashaAggregator的策略是:仅当_gradient_estimator is None(即首轮)时,通过配置项SEND_FULL_GRADIENT=True要求客户端回传未压缩的完整梯度;其余轮次全部走压缩通道。而MarinaAggregator不同——它按概率p = 压缩向量大小 / 参数维度做伯努利采样,随机要求客户端在某些轮次回传完整梯度(对应 MARINA 算法中c_k的随机切换),因此 MARINA 在部分轮次仍需传输完整向量,这正是 DASHA "只发压缩向量"更实用的原因。

客户端:梯度计算与压缩(client.py)

client.py 实现了客户端逻辑:

  • CompressionClient:抽象基类,负责参数同步(将一维参数向量 reshape 回各层)、压缩器维度设置;
  • DashaClient:确定性 DASHA 客户端。首轮计算完整梯度并初始化局部/全局梯度估计器;后续轮次按论文 Algorithm 1 的第 8、9 行,压缩g_i - ĝ_i - momentum·(ĝ - ĝ_i)这一差分项,其中动量momentum = 1 / (1 + 2·ω)(ω 为压缩器方差,取自论文 Theorem 6.1),压缩后更新本地梯度估计器;
  • MarinaClient:首轮回传完整梯度;后续轮次压缩g_i - 上次梯度的差分;
  • StochasticDashaClient/StochasticMarinaClient:随机化变体,基于小批量采样计算随机梯度,其中_calculate_stochastic_gradient_in_current_and_previous_parameters会在当前与上一组参数上分别计算梯度,以实现随机方差缩减。

压缩器(compressors.py)

compressors.py 定义了论文使用的 RandK 稀疏化压缩器:

  • RandKCompressor:从向量中无放回随机选取 K 个坐标,将选中坐标值乘以dim / K作为无偏缩放,其余坐标置零;其方差ω = dim / K - 1。K 由配置项compressor.number_of_coordinates控制;
  • IdentityUnbiasedCompressor:恒等(不压缩)压缩器,用于首轮回传完整梯度,方差为 0;
  • decompress:按索引将压缩向量还原为稠密向量;estimate_size:估算压缩向量占用比特数(索引与值的位数之和),供绘图脚本绘制"横轴为每位客户端通信比特数"的收敛曲线。

小规模实验:mushrooms 上对比 DASHA 与 MARINA

下面的命令会同时运行 DASHA 与 MARINA,遍历不同的step_size,其余参数与论文一致。同时设置method.client.send_gradient=true,让客户端回传完整梯度,以便服务端计算梯度范数(squared_gradient_norm)这一收敛性指标。

# Run experiments python -m dasha.main --multirun method=dasha,marina compressor.number_of_coordinates=10 method.strategy.step_size=0.25,0.5,1.0 method.client.send_gradient=true # The previous script output paths to the results (ex: multirun/2023-09-16/10-39-30/1 multirun/2023-09-16/10-39-30/2 ...). # Plot results python -m dasha.plot --input_paths multirun/2023-09-16/10-39-30/1 multirun/2023-09-16/10-39-30/2 --output_path plot.png --metric squared_gradient_norm # or it is sufficient to give the common folder as input python -m dasha.plot --input_paths multirun/2023-09-16/10-39-30 --output_path plot.png --metric squared_gradient_norm

Hydra 的--multirun会将每次运行的结果保存到multirun/<日期>/<时间>/<job_id>/目录下,每个目录内包含config.yaml(本次运行完整配置)与history(Flower History 对象,含分布式指标)。绘图脚本 plot.py 支持以下参数:

参数含义默认值
--input_paths结果目录(可多个,或直接给公共父目录)必填
--output_path输出图片路径必填
--metric绘制的指标:loss/squared_gradient_norm/accuracyloss
--smooth-plot滑动平均窗口大小(大模型曲线噪声大时建议设置,如 100)None

绘图脚本横轴统一为"每位客户端收到的比特数"(#bits / client,取自服务端记录的received_bytes指标),纵轴为所选指标,从而直观对比两种方法在相同通信预算下的收敛速度;纵轴对数刻度(loss 与 squared_gradient_norm 会自动启用 log 刻度)。

上述命令生成的结果与下图类似:

小规模实验:DASHA 与 MARINA 收敛对比(mushrooms)

大规模实验:CIFAR10 上的 ResNet18 训练

以下实验在 CIFAR10 数据集上对比 DASHA 与 MARINA 训练 ResNet18(含 logistic 损失)。由于模型参数量大,这里将压缩坐标数提升到 200 万,并启用 GPU 与精度评估:

# Run experiments python -m dasha.main method.strategy.step_size=0.01 method=stochastic_dasha num_rounds=10000 compressor.number_of_coordinates=2000000 model=resnet_18_with_logistic_loss method.client.strict_load=false dataset=cifar10 method.client.device=cuda method.client.evaluate_accuracy=true local_address=localhost:8001 method.client.mega_batch_size=16 python -m dasha.main method=stochastic_marina method.strategy.step_size=0.01 num_rounds=10000 compressor.number_of_coordinates=2000000 model=resnet_18_with_logistic_loss method.client.strict_load=false dataset=cifar10 method.client.device=cuda method.client.evaluate_accuracy=true local_address=localhost:8002 # The previous scripts output paths to the results. We define them as PATH_DASHA and PATH_MARINA # Plot results python -m dasha.plot --input_paths PATH_DASHA PATH_MARINA --output_path plot_nn.png --smooth-plot 100

这里需要注意几点:

  • 两条命令使用不同的local_addresslocalhost:8001localhost:8002),避免并行运行时端口冲突;
  • method.client.strict_load=false是因为不同客户端加载同一 ResNet 结构时可能出现参数键不完全匹配的情况,放宽校验以保证运行;
  • 随机化变体(stochastic_*)要求显式给定method.strategy.step_size(配置中为???必填项);
  • --smooth-plot 100对 10000 轮的噪声曲线做窗口为 100 的滑动平均,使对比更清晰。

预期生成的大规模实验对比图如下:

大规模实验:CIFAR10 上 DASHA 与 MARINA 训练 ResNet18 对比

运行测试

Baseline 自带单元测试与集成测试,测试代码位于 dasha/tests:

# Run unit tests pytest ./dasha/tests/ # Run unit and integration tests. Some long integration tests are turned off be default. TEST_DASHA_LEVEL=1 pytest ./dasha/tests/

测试覆盖了客户端逻辑(test_clients.py)、压缩器正确性(test_compressors.py)、基线整体流程(test_dasha_baseline.py)、数据集加载(test_datasets.py)与模型(test_models.py)。其中部分耗时的集成测试默认关闭,需要设置环境变量TEST_DASHA_LEVEL=1才会执行。

总结

通过本 baseline,你可以在 Flower 框架下完整复现 DASHA 论文的核心实验:小规模场景下,在 mushrooms 数据集上以squared_gradient_norm为指标对比 DASHA 与 MARINA 在不同步长下的收敛曲线;大规模场景下,在 CIFAR10 上以 ResNet18 验证随机化变体(stochastic DASHA / stochastic MARINA)在通信受限条件下的表现。从源码层面看,DASHA 相比 MARINA 的实践优势(只传压缩向量)体现在服务端MarinaAggregator需要按概率随机回传完整梯度、而DashaAggregator仅首轮需要完整梯度这一关键差异上;RandK 压缩器、方差缩减动量与"横轴为通信比特数"的对比绘图方式,共同构成了这套可直接复用的压缩通信实验范式。

【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询