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 |
ResNet18WithLogisticLoss | ResNet18 网络 + 交叉熵损失(论文 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个客户端:
| 数据集 | 类别数 | 划分方式 |
|---|---|---|
| mushrooms | 2 | random |
| CIFAR10 | 10 | random |
数据集的加载逻辑在 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-core、scikit-learn、matplotlib与torch/torchvision。PyTorch 的安装方式因平台而异:Linux 默认安装 CUDA 11.8 版(torch-2.0.0+cu118),macOS 则安装 CPU 版,若环境不同请相应修改 pyproject.toml 中[tool.poetry.dependencies]的torch与torchvision行。
运行实验
激活 Poetry 环境(在 baselines/dasha 目录下执行poetry shell)后,即可运行默认配置:
python -m dasha.main # this will run using the default settings in `dasha/conf`默认配置定义在 conf/base.yaml 中:num_clients: 5、num_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_coordinates | RandK 压缩器每轮保留的坐标数 K | 1 |
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=marinamethod配置组(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_normHydra 的--multirun会将每次运行的结果保存到multirun/<日期>/<时间>/<job_id>/目录下,每个目录内包含config.yaml(本次运行完整配置)与history(Flower History 对象,含分布式指标)。绘图脚本 plot.py 支持以下参数:
| 参数 | 含义 | 默认值 |
|---|---|---|
--input_paths | 结果目录(可多个,或直接给公共父目录) | 必填 |
--output_path | 输出图片路径 | 必填 |
--metric | 绘制的指标:loss/squared_gradient_norm/accuracy | loss |
--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_address(localhost:8001与localhost: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),仅供参考