Flower 安全聚合协议解析:SecAgg 与 SecAgg+ 在联邦学习中的实现
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
联邦学习中,服务器在每一轮聚合后能看到各客户端上传的模型更新,而模型更新本身可能泄露客户端的本地数据。安全聚合(Secure Aggregation)协议正是为解决这一问题而设计:在将模型更新发送给服务器之前对其进行加密掩码,使得服务器只能解密出"聚合后的模型更新",而无法窥探任何单个客户端的贡献。本文基于 Flower 框架的官方解释文档 explanation-ref-secure-aggregation-protocols.rst 展开,结合 SecAggPlusWorkflow 与 secaggplus_mod 的源码实现,完整讲解 Flower 中 SecAgg / SecAgg+ 协议的工作原理、参数体系与可运行的配置方式。读完后,你将能够:理解 SecAgg+ 的四阶段协议流程与密钥共享机制;掌握SecAggPlusWorkflow各参数的含义、取值约束与调参权衡;复现 flower-secure-aggregation 示例并为其配置协议参数。
什么是安全聚合协议
安全聚合协议指的是SecAgg、SecAgg+、LightSecAgg、FastSecAgg等一系列协议的统称,该概念最早由 Bonawitz 等人在论文Practical Secure Aggregation for Federated Learning on User-Held Data中提出。其核心目标是:
- 保密性:服务器只能获取所有客户端更新的加权和,任何单个客户端的更新(包括权重因子)均不可见;
- 容错性:允许部分客户端在聚合过程中掉线,只要密钥份额满足重构阈值,仍可恢复其贡献并保证聚合正确性;
- 权重因子保护:每个客户端上传的不只是模型参数
params,还有权重因子w(即本地样本数num_examples),两者均被打上隐私掩码。
Flower 目前内置实现了SecAgg和SecAgg+两个协议;从源码结构看,SecAgg是SecAgg+的一个特例——SecAggWorkflow 直接继承自SecAggPlusWorkflow,二者的区别在于密钥份额的划分方式:SecAgg 中每个客户端的私钥被拆分为 N 份(N 为被选中的客户端总数),而 SecAgg+ 中份数num_shares可独立配置,通常小于 N,从而降低通信开销。除内置协议外,官方文档也指出,未来会实现更多协议,用户也可以基于 Flower 的低层 API 自行实现自定义的安全聚合协议。
上图是官方文档中给出的SecAgg+协议时序图:虚线代表跨网络通信,实线代表同一进程内的调用。ServerApp连接到SuperLink,ClientApp连接到SuperNode,因此服务器端与客户端之间的所有消息都经由 SuperLink 和 SuperNode 转发——这也意味着安全聚合的"邻居转发"并不直接发生在客户端之间,而是由服务器代为转发加密密文。
Flower 中的组件映射:ServerApp、ClientApp 与 Mod
在 Flower 的新式 App 架构中,安全聚合被拆分为两个正交的组件:
| 角色 | 组件 | 源码位置 | 职责 |
|---|---|---|---|
| 服务端 | SecAggPlusWorkflow | secaggplus_workflow.py | 作为 fit workflow 编排四个协议阶段,执行聚合与去掩码 |
| 客户端 | secaggplus_mod | secaggplus_mod.py | 作为 ClientApp 的 mod,拦截训练消息、执行密钥生成、掩码计算与转发 |
| 服务端(简化) | SecAggWorkflow | secagg_workflow.py | SecAgg+ 的特例,份额数固定为被选中客户端数 |
| 客户端(简化) | secagg_mod | secagg_mod.py | 对应 SecAgg 的客户端 mod |
两者的组合方式是声明式的:ServerApp的main()中以SecAggPlusWorkflow构造DefaultWorkflow的 fit 工作流,ClientApp通过mods=[secaggplus_mod]注册 mod。一个完整的可运行示例见 examples/flower-secure-aggregation 目录,其 server_app.py 与 client_app.py 展示了完整接线方式,下文将逐段解读。
SecAggPlusWorkflow:参数体系与取值约束
SecAggPlusWorkflow的构造函数是协议行为的唯一配置入口。以下参数说明直接来自 源码 docstring,并结合_check_init_params的校验逻辑给出取值约束:
fit_workflow = SecAggPlusWorkflow( num_shares=3, # 必需:私钥份额数(int 或 float) reconstruction_threshold=2, # 必需:重构私钥所需的最小份额数(int 或 float) max_weight=1000.0, # 可选,默认 1000.0:单个客户端权重的上界 clipping_range=8.0, # 可选,默认 8.0:量化前的裁剪范围 quantization_range=4194304, # 可选,默认 2**22:量化目标范围 modulus_range=4294967296, # 可选,默认 2**32:掩码取模范围 timeout=None, # 可选,默认 None:每阶段等待回复的超时(秒) )num_shares:每个客户端私钥被拆成的份额数。取 int 时必须大于 2(num_shares == 1会被自动转换为比例1.0);取 float 时表示"被选中客户端总数"的比例,实际份额数在运行时动态计算(见 setup_stage:round(num_shares * num_samples),并保证不超过样本数、不小于 3)。份额数越大,对掉线越鲁棒,但通信与计算开销也越高;源码中还会在份额数为偶数时发出警告并自动 +1("The number of shares in the SecAgg+ protocol should be odd")。
reconstruction_threshold:重构某客户端私钥所需的最小份额数,决定隐私强度与容错上限。取 float 时必须落在 (0, 1],表示份额总数的比例,且运行时至少为 2。阈值越高,隐私保证越强(任意低于阈值的客户端子集都无法恢复单个客户端的秘密),但可容忍的掉线也越少。校验逻辑见 _check_init_params:int 时必须小于num_shares。
max_weight:单个客户端在加权平均中可贡献的最大权重(即最大样本数)。默认 1000.0。注意官方注释提醒:max_weight过大会损害量化精度,因为它与clipping_range、quantization_range共同决定量化后的取值上界。
clipping_range / quantization_range / modulus_range:三个参数共同构成"裁剪 → 量化 → 模运算"的数值化管线,把浮点模型更新转换为可执行密码学操作的整数向量:
clipping_range(默认 8.0):参数先被裁剪到[-clipping_range, clipping_range];quantization_range(默认2**22 = 4194304):裁剪后的浮点数被量化到[0, quantization_range - 1]的整数;modulus_range(默认2**32 = 4294967296):掩码元素在此范围内均匀采样,且所有加法聚合都在该模下进行。源码校验它必须是 2 的幂(通过二进制中 "1" 的个数判断,见 L256-L265),并且必须严格大于quantization_range,以避免溢出。
timeout:每个send_and_receive阶段等待客户端回复的秒数;为None时无限等待。配合num_shares/reconstruction_threshold,超时导致客户端掉线时,只要份额仍满足阈值,聚合可以继续。
四阶段协议流程:从 setup 到 unmask
SecAggPlusWorkflow.__call__将协议拆成四个串行阶段(L183-L202):setup_stage → share_keys_stage → collect_masked_vectors_stage → unmask_stage,任一阶段返回False(如活跃邻居数不足阈值)即中止本轮聚合。各阶段职责与源码对应如下:
阶段 0:setup —— 下发配置并收集公钥
setup_stage 首先调用strategy.configure_fit完成客户端采样(因此采样仍由你使用的 Strategy,如 FedAvg 控制),随后:
- 根据采样数动态确定
num_shares与threshold; - 随机打乱节点 ID,为每个节点构建环形邻居关系:
nid_to_neighbours中每个节点与其前后各num_shares // 2个节点互为邻居(L348-L358)——邻居即"密钥份额的持有者"; - 将协议配置(
STAGE=SETUP、样本数、份额数、阈值、裁剪/量化/模范围、max_weight)打包为ConfigRecord下发; - 收集各客户端回复的两对非对称公钥
pk1、pk2(客户端侧由secaggplus_mod通过generate_key_pairs生成)。
每个阶段结束都会执行_check_threshold:任一采样节点的活跃邻居数低于阈值即判定失败(L267-L273)。
阶段 1:share keys —— 广播公钥、收集加密密钥份额
share_keys_stage(L402-L470)将每个节点的邻居们的公钥广播给该节点。客户端收到后:对自己的私钥执行 Shamir 秘密分享(create_shares),把各份额用对应邻居的pk1加密后,连同DESTINATION_LIST与CIPHERTEXT_LIST一并回复。服务器在这里只做"密文交换机":它解析每个回复中的源节点、目标节点与密文三元组,重建出forward_ciphertexts(目标节点 → 密文列表)与forward_srcs(目标节点 → 来源列表)。由于所有份额均以邻居公钥加密,服务器转发的只是不可解密的密文,看不到任何密钥内容。
阶段 2:collect masked vectors —— 转发密文并收集掩码向量
collect_masked_vectors_stage(L472-L541)把上一阶段收到的密文份额转发给目标客户端,同时下发FitIns(训练指令)。客户端侧secaggplus_mod此时:用收到的密文份额与自己的sk1恢复出共享秘密,生成本地的"私有掩码"(PRG 种子),并对[w, w * params](权重因子与加权参数)执行裁剪、量化后加上私有掩码及成对掩码,上传MASKED_PARAMETERS。服务器将收到的掩码向量逐元素相加并做模运算(parameters_addition+parameters_mod,L516-L530),得到全量掩码和。由于掩码在模下可加,各客户端的私有掩码与成对掩码在求和后仍保持可被后续阶段消除的性质。
阶段 3:unmask —— 收集份额、去掩码并聚合
unmask_stage(L543-L677)是隐私保护的收口:
- 服务器向各活跃客户端通报其邻居中的活跃/掉线名单;
- 客户端回复自己持有的密钥份额(
NODE_ID_LIST+SHARE_LIST); - 服务器对每个采样节点用
combine_shares做 Shamir 重构:- 活跃节点:重构出其 PRG 种子,重新生成私有掩码并从掩码和中减去;
- 掉线节点:重构出其
sk1,再与其各邻居的pk1生成共享密钥、重算成对掩码,按节点 ID 大小关系执行加/减,从而"移除"掉线客户端的贡献,保证聚合结果只覆盖实际参与者;
- 经
factor_extract+dequantize反量化,并叠加偏移-(len(active_nids) - 1) * clipping_range抵消量化引入的系统偏差,还原为浮点聚合参数; - 最后将结果回填到
legacy_results中的FitRes.parameters,交给Strategy.aggregate_fit完成最终聚合(L665-L668)——也就是说,策略层拿到的已经是隐私保护的聚合结果,Strategy 本身无需感知安全聚合的存在。
客户端侧SecAggPlusState(L64-L120)保存了协议所需的全部秘密:两对公私钥sk1/pk1、sk2/pk2、PRG 随机种子rd_seed、各密钥份额字典等,其序列化/反序列化(to_dict及对应解析)保证了状态可以在 ClientApp 的消息间正确传递。
实战:运行 flower-secure-aggregation 示例
examples/flower-secure-aggregation 是官方给出的 SecAgg+ 完整示例,它把 quickstart-pytorch 同样的 PyTorch 训练负载换成了安全聚合。目录结构:
flower-secure-aggregation ├── secaggexample │ ├── client_app.py # 定义 ClientApp,注册 secaggplus_mod │ ├── server_app.py # 定义 ServerApp,构造 SecAggPlusWorkflow │ ├── task.py # 模型、训练与数据加载 │ └── workflow_with_log.py # is-demo=true 时使用的带日志工作流 ├── pyproject.toml # 依赖与默认配置 └── README.mdclient_app.py 中,安全聚合能力完全由一行 mod 声明注入:
app = ClientApp( client_fn=client_fn, mods=[ secaggplus_mod, ], )server_app.py 中,协议参数从context.run_config读取并传入 workflow:
strategy = FedAvg( fraction_fit=1.0, min_fit_clients=5, fraction_evaluate=(0.0 if is_demo else context.run_config["fraction-evaluate"]), min_available_clients=5, evaluate_metrics_aggregation_fn=weighted_average, initial_parameters=parameters, ) context = LegacyContext( context=context, config=ServerConfig(num_rounds=num_rounds), strategy=strategy, ) fit_workflow = SecAggPlusWorkflow( num_shares=context.run_config["num-shares"], reconstruction_threshold=context.run_config["reconstruction-threshold"], max_weight=context.run_config["max-weight"], ) workflow = DefaultWorkflow(fit_workflow=fit_workflow) workflow(grid, context)示例的默认参数在 pyproject.toml 中:
[tool.flwr.app.config] num-server-rounds = 3 fraction-evaluate = 0.5 local-epochs = 1 learning-rate = 0.1 batch-size = 32 # SecAgg+ 协议参数 num-shares = 3 reconstruction-threshold = 2 max-weight = 9000 timeout = 120.0 is-demo = true注意示例把max-weight设为 9000(CIFAR-10 分片样本数较多时的上界),并遵循"max_weight 过大会损害量化精度"的提示,在演示模式下改用max_weight=1。运行方式:
# 安装框架 pip install flwr # 获取示例后安装 pip install -e . # 以 Simulation Engine 运行(默认模式) flwr run . --stream # 覆盖配置,例如减少轮数、调整学习率 flwr run . --run-config "num-server-rounds=5 learning-rate=0.25" --stream # 切换为真实训练(关闭演示模式) flwr run . --run-config is-demo=false --stream演示模式(is-demo=true)额外启用了两个特性:一是使用SecAggPlusWorkflowWithLogs(workflow_with_log.py)打印各阶段调试日志;二是客户端支持drop指令与timeout延迟,用于模拟客户端掉线、验证阈值容错——这恰好对应num_shares=3、reconstruction_threshold=2的配置含义:任意 1 个客户端掉线,其余邻居仍持有 ≥2 份份额,聚合可继续。切换到部署模式(Deployment Engine)时,同样的应用代码无需修改,即可跑在真实的 SuperLink/SuperNode 拓扑上,此时建议再叠加 TLS 与 SuperNode 认证以保护传输链路。
参数调优建议与工程要点
综合源码中的校验逻辑与 docstring 注释,实际部署时可以从三个维度权衡:
- 隐私强度 vs 容错:
reconstruction_threshold越接近num_shares,越难以被少于阈值的合谋集合攻破,但可容忍的掉线越少。默认示例的 (3, 2) 组合是"低开销、可容忍 1 个掉线"的起点,生产环境可提升到 (5, 3) 或更高。 - 数值精度 vs 溢出:
modulus_range必须是 2 的幂且大于quantization_range;num_shares过大时源码会警告可能触发模运算溢出(num_shares > modulus_range / quantization_range时)。量化参数三件套(clipping_range、quantization_range、modulus_range)需要与模型参数分布匹配:裁剪范围过窄会引入信息损失,量化范围过小则精度下降。 - 权重约束:
max_weight必须覆盖实际客户端样本数的上界,否则会截断权重因子;同时它是精度敏感参数,建议尽量贴近真实样本规模而不是取很大的冗余值。
最后需要明确适用边界:SecAggPlusWorkflow要求传入LegacyContext(否则抛出TypeError,见 L183-L188),即它工作在 Flower 兼容传统 Strategy 体系的 workflow 层;策略(FedAvg 等)仍负责采样与最终聚合,安全聚合只替换"收集参数 → 聚合"这一段的可见性。若你的场景需要LightSecAgg、FastSecAgg等协议,目前框架尚未内置,需要基于低层 API 自行实现。
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考