ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

Flower Baseline 实践:在 Flower 中复现 DASHA 分布式非凸优化与通信压缩算法 Flower Baseline 实践在 Flower 中复现 DASHA 分布式非凸优化与通信压缩算法【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flowerDASHADistributed 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 方法 MARINAGorbunov et al., 2020DASHA 系列在理论上改进了 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.yamlResNet18WithLogisticLossResNet18 网络 交叉熵损失论文 Section A.4conf/model/resnet_18_with_logistic_loss.yaml值得说明的是论文中使用的非凸损失NonConvexLoss在 models.py 中有明确实现它将目标标签映射到{-1, 1}计算sigmoid(w·x·y)后取(1 - sigmoid)^2的均值这一设计正是为了在简单模型上验证算法处理非凸目标的收敛行为。数据集与划分方式默认数据集按**随机划分random**方式切分给n个客户端数据集类别数划分方式mushrooms2randomCIFAR1010random数据集的加载逻辑在 dataset.py 中mushrooms 通过sklearn.datasets.load_svmlight_file读取 LIBSVM 格式文件并将标签重映射为 0/1CIFAR10 使用 Torchvision 加载并做ToTensor Normalize预处理。切分则使用torch.utils.data.random_split将整个训练集按1/num_clients等分给各客户端dataset.py 的random_split。若未指定path_to_datasetdataset_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.0cu118macOS 则安装 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默认数据集为 libsvmmushrooms、默认方法为 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_size0.5 compressor.number_of_coordinates10 # 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.devicecuda常用配置项说明配置路径含义默认值method.strategy.step_size服务端聚合时的更新步长dasha/marina 为 0.5stochastic 变体为必填???compressor.number_of_coordinatesRandK 压缩器每轮保留的坐标数 K1method.client.device客户端训练设备cpumethod.client.send_gradient是否让客户端在 evaluate 阶段回传完整梯度用于计算梯度范数指标falsemethod.client.mega_batch_size随机化变体的 mega-batch 大小计算初始梯度时用dasha: 100marina: 10method.client.batch_size随机化变体采样的小批量大小25method.client.strict_load加载服务端参数时是否严格要求无缺失键truedatasetcifar10切换数据集libsvmmushroomsnum_rounds联邦训练轮数10000local_address服务端监听地址多进程并行运行时使用localhost:8080运行 MARINA 对照论文以 MARINAGorbunov et al., 2020为对照基线切换方法只需一行python -m dasha.main methodmarinamethod配置组conf/method内置了四种方法dasha确定性 DASHAdasha.yaml客户端为DashaClientmarina确定性 MARINAmarina.yaml客户端为MarinaClientstochastic_dasha随机化 DASHAstochastic_dasha.yaml客户端为StochasticDashaClientstochastic_marina随机化 MARINAstochastic_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 1MARINA 的逻辑几乎相同服务端维护全局参数_parameters与梯度估计器_gradient_estimator每轮收集各客户端返回的压缩向量先估算每个客户端收到的比特数estimate_size再解压并取平均若_gradient_estimator为 None首轮则直接将该均值作为初始梯度估计器否则累加最后执行_parameters - step_size * gradient_estimator完成一步更新。DashaAggregator的策略是仅当_gradient_estimator is None即首轮时通过配置项SEND_FULL_GRADIENTTrue要求客户端回传未压缩的完整梯度其余轮次全部走压缩通道。而MarinaAggregator不同——它按概率p 压缩向量大小 / 参数维度做伯努利采样随机要求客户端在某些轮次回传完整梯度对应 MARINA 算法中c_k的随机切换因此 MARINA 在部分轮次仍需传输完整向量这正是 DASHA 只发压缩向量更实用的原因。客户端梯度计算与压缩client.pyclient.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.pycompressors.py 定义了论文使用的 RandK 稀疏化压缩器RandKCompressor从向量中无放回随机选取 K 个坐标将选中坐标值乘以dim / K作为无偏缩放其余坐标置零其方差ω dim / K - 1。K 由配置项compressor.number_of_coordinates控制IdentityUnbiasedCompressor恒等不压缩压缩器用于首轮回传完整梯度方差为 0decompress按索引将压缩向量还原为稠密向量estimate_size估算压缩向量占用比特数索引与值的位数之和供绘图脚本绘制横轴为每位客户端通信比特数的收敛曲线。小规模实验mushrooms 上对比 DASHA 与 MARINA下面的命令会同时运行 DASHA 与 MARINA遍历不同的step_size其余参数与论文一致。同时设置method.client.send_gradienttrue让客户端回传完整梯度以便服务端计算梯度范数squared_gradient_norm这一收敛性指标。# Run experiments python -m dasha.main --multirun methoddasha,marina compressor.number_of_coordinates10 method.strategy.step_size0.25,0.5,1.0 method.client.send_gradienttrue # 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本次运行完整配置与historyFlower History 对象含分布式指标。绘图脚本 plot.py 支持以下参数参数含义默认值--input_paths结果目录可多个或直接给公共父目录必填--output_path输出图片路径必填--metric绘制的指标loss/squared_gradient_norm/accuracyloss--smooth-plot滑动平均窗口大小大模型曲线噪声大时建议设置如 100None绘图脚本横轴统一为每位客户端收到的比特数#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_size0.01 methodstochastic_dasha num_rounds10000 compressor.number_of_coordinates2000000 modelresnet_18_with_logistic_loss method.client.strict_loadfalse datasetcifar10 method.client.devicecuda method.client.evaluate_accuracytrue local_addresslocalhost:8001 method.client.mega_batch_size16 python -m dasha.main methodstochastic_marina method.strategy.step_size0.01 num_rounds10000 compressor.number_of_coordinates2000000 modelresnet_18_with_logistic_loss method.client.strict_loadfalse datasetcifar10 method.client.devicecuda method.client.evaluate_accuracytrue local_addresslocalhost: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:8001与localhost:8002避免并行运行时端口冲突method.client.strict_loadfalse是因为不同客户端加载同一 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_LEVEL1 pytest ./dasha/tests/测试覆盖了客户端逻辑test_clients.py、压缩器正确性test_compressors.py、基线整体流程test_dasha_baseline.py、数据集加载test_datasets.py与模型test_models.py。其中部分耗时的集成测试默认关闭需要设置环境变量TEST_DASHA_LEVEL1才会执行。总结通过本 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),仅供参考
返回列表