简介:面向联邦学习安全聚合研究的一套可运行代码实现,重点给出基于Shamir门限秘密共享的FedSTSS模型,并配套FedShare、Scotch、FedAvg等基线方法的对比实验。代码包含服务端与客户端Python脚本、秘密共享与模型聚合核心模块、多数据集加载处理逻辑,以及一键启动/清理的Shell脚本,适合作为毕设、课设或安全聚合方向入门进阶的参考工程。整个压缩包共54个文件,以py源码为主(22个),辅以sh脚本、运行日志、CSV数据与说明文档,包体仅68KB,轻量易部署。已有140人学习浏览,项目说明中标注代码均通过测试、答辩评分96分,并支持远程教学,可信度较高。下载后可核对目录结构,按需运行FedSTSS等对比实验,也可基于现有模块扩展隐私保护机制或更换数据集进行二次开发。
1. FedSTSS在解决什么问题:联邦学习里的明文梯度与门限秘密共享
普通FedAvg跑起来之后,服务器每轮都能拿到参与方的明文梯度或模型更新。过去几年已经有不少工作证明,这些更新携带训练样本的大量信息,攻击者不需要侵入客户端,只用服务器侧拿到的梯度就能反推出原始图片和文本。数据不出本地,并不代表数据真的安全。联邦学习要落地到医疗、金融这类敏感场景,安全聚合必须是第一层防护而不是可选加分项。
目前主流的安全聚合路线有三类:同态加密、差分隐私、基于Shamir门限秘密共享的掩码方案。FedSTSS走的是第三条路。核心思路是每个客户端在提交更新前先加一个随机掩码,再把这个掩码用Shamir(t,n)门限方案切成多份分发出去。服务器只有在收集到足够多份额、恢复出掩码总和之后,才能解开这一轮所有客户端更新的总和——但始终看不到任何单个客户端的明文更新。
这篇文章从Shamir的有限域实现讲起,然后给出一套可运行的Python源码结构和对比实验设计。适用对象是正在搭建横向联邦学习框架的工程师,以及需要为课程项目或论文补一组安全聚合实验的同学。下面从原理开始,逐步落到代码和参数。
2. Shamir门限秘密共享原理与FedSTSS聚合流程
2.1 拉格朗日插值与Shamir(t,n)的份额生成
Shamir门限方案的核心是这样一个事实:任意t个平面上的点能唯一确定一个t-1次多项式。要把秘密整数s切分成n份,就构造一个t-1次多项式:
f(x) = s + a₁x + a₂x² + … + a₍ₜ₋₁₎x⁽ᵗ⁻¹⁾ mod p
其中p是一个大素数,系数a₁到a₍ₜ₋₁₎是在[0, p)上随机选择的整数。第i个参与方拿到的份额是点(i, f(i)),i从1取到n。因为多项式在模p的有限域上定义,少于t个份额时,候选秘密s可以对应无穷多个不同的多项式,攻击者在信息论意义上无法区分出真实秘密;只有拿到任意t个份额,才能通过拉格朗日插值计算出f(0)得到s。
重构公式在实现时需要写成有限域运算。给定t个份额点(xᵢ, yᵢ),秘密为:
s = Σ yᵢ · Πⱼ≠ᵢ (0 - xⱼ) / (xᵢ - xⱼ) mod p
注意这里的除法是模逆运算,不是浮点除法。Python 3.8及以上可以用pow(den, -1, p)直接求模逆,这也是我会在代码里依赖的运行时版本。
2.2 FedSTSS如何用份额恢复掩码总和
FedSTSS的关键设计是利用Shamir方案的线性同态性质。如果不同客户端分别用自己的掩码m₁、m₂构造了多项式f₁(x)、f₂(x),那么把它们逐项相加得到的新多项式h(x) = f₁(x) + f₂(x)仍然是一个t-1次多项式,并且h(0) = m₁ + m₂。这意味着对于相同的x坐标,把多个份额相加后再插值,可以直接恢复多个掩码的总和,而不需要暴露任何一个单独的掩码。
具体到一轮联邦训练,流程分五步。第一步,每个客户端在本地训练得到更新量Δw。第二步,客户端生成与Δw同形状的随机掩码向量m,计算带掩码的更新c = Δw + m,这里所有运算都在模p的整数域上完成。第三步,客户端把掩码向量m的每个元素分别执行Shamir(t,n)分割,得到n份份额。第四步,客户端把自己的第i份掩码份额发送给参与方i;每个参与方把收到的同一x坐标上的份额相加,得到一条"汇总份额"。第五步,服务器从在线客户端收集至少t条汇总份额,先用拉格朗日插值恢复出所有客户端掩码的总和Σm,然后计算Σc - Σm得到明文更新总和,最后除以参与方数量完成平均。
这里要特别注意:服务器拿到的是每个客户端掩码的至多一个坐标份额,而不是某个客户端的t份份额。所以服务器能恢复掩码总和,却无法恢复任何一个客户端的独立掩码。这是我判断FedSTSS实现是否安全的关键检查点。
2.3 为什么不优先选同态加密或差分隐私
同态加密在联邦学习里最常见的是Paillier这种加法同态方案。它的问题不是安全性,而是性能和密钥管理。密文域上的加法比明文慢两个数量级,密文长度也会扩张到256位甚至更长。参与方上百、模型参数上百万时,单轮通信和CPU开销都很难接受。差分隐私则是另一种取舍:它通过在梯度上添加噪声保护隐私,噪声越大隐私越强,但模型精度必然下降,而且隐私预算的消耗是逐轮累积的。
Shamir门限方案的优势在于计算廉价,只涉及有限域上的加法和乘法,不涉及幂运算或格运算;同时它不向聚合结果注入噪声,理论上不损失精度。门限t还天然对应容错:允许掉线的参与方数量是n - t,这在真实联邦环境里非常有价值。下表是三种方案的对比:
| 方案 | 梯度隐私 | 单轮计算开销 | 通信扩张倍数 | 掉线容错 | Python实现难度 |
|---|---|---|---|---|---|
| FedAvg | 无 | 低 | 1x | 强 | 低 |
| Paillier同态加密 | 有 | 高 | 8~16x | 弱 | 高 |
| Shamir门限(FedSTSS) | 有 | 低 | 2t/n 至 2x | 可配置 | 中 |
FedSTSS的代价主要是通信和协调复杂度,掩码份额的交换在参与方数量大时会有O(n²)的传输压力。工程上一般会配合伪随机数生成器来压缩份额,或者把份额交换合并到聚合通道里做批处理,我后面会提到这一类优化方向。
2.4 向量的掩码处理与量化前提
联邦学习里的更新是浮点向量,而Shamir运算必须在整数有限域上操作,所以FedSTSS的第一步量化编码不能省。常见做法是给梯度乘以一个缩放因子scale后取整,再模p;解码时把超过p/2的值减掉p还原成负数。scale的选择直接影响聚合精度和溢出风险:scale太小量化误差大,scale太大会让累加和接近模数p导致截断错误。一般先估算梯度绝对值的上界、客户端数量和联邦轮次,再倒推scale。比如梯度绝对值上限1.0,客户端10个,那么scale取10000时累加和绝对值的上界是10万,远小于2⁶¹这样的大素数,余量充足。
3. FedSTSS的Python实现:从秘密分割到联邦安全聚合的最小工程
3.1 目录结构与依赖说明
一个能跑通的最小FedSTSS工程通常只需要五个文件。依赖方面用Python 3.8以上版本、numpy,对比实验部分可选scikit-learn和phe库。目录结构如下:
| 文件 | 职责 |
|---|---|
| shamir.py | 有限域上的份额生成、拉格朗日重构 |
| quant.py | 浮点梯度与有限域整数之间的编码解码 |
| client_side.py | 客户端本地训练、掩码生成与份额交换逻辑 |
| server_side.py | 服务器聚合、掩码总和恢复与参数更新 |
| run_comparison.py | 对比实验入口,支持fedavg、fedstss、paillier三种模式 |
安装依赖时用pip install numpy scikit-learn phe即可,phe只在跑Paillier对照实验时才需要。下面先实现底层 shamir.py,它是整个工程正确性的根基。
3.2 有限域上的Shamir分割与重构实现
# shamir.py import random def generate_shares(secret, threshold, num_shares, prime): """把秘密 secret 切分成 num_shares 份,至少 threshold 份可恢复""" coeffs = [secret % prime] + [ random.randrange(prime) for _ in range(threshold - 1) ] shares = [] for x in range(1, num_shares + 1): # 霍纳法计算多项式值,避免每次迭代都做幂运算 y = 0 for c in reversed(coeffs): y = (y * x + c) % prime shares.append((x, y)) return shares def reconstruct_secret(share_points, prime): """拉格朗日插值求 f(0),share_points 长度必须 >= threshold""" secret = 0 for i, (xi, yi) in enumerate(share_points): num = 1 den = 1 for j, (xj, _) in enumerate(share_points): if i == j: continue num = (num * xj) % prime # 分子:连乘 0 - xj 在模意义下即 -xj den = (den * (xj - xi)) % prime # 分母:连乘 (xi - xj) l_i = num * pow(den, -1, prime) % prime secret = (secret + yi * l_i) % prime return secret这段代码有两个点需要解释。第一,生成份额时用霍纳法而不是逐项计算x的幂,时间复杂度从O(t²)降到O(t),在t较大时差距明显。第二,重构时分子直接写成xj,是因为f(0)代入拉格朗日基函数后,分子是(0 - xj)的连乘,在模p下等于(-xj) mod p,等价于p - xj,而代码里用xj再在最后取模会有符号问题……这里更稳妥的写法是num = (num * (p - xj)) % prime,或者保持num = (num * (-xj)) % prime。上面的写法num = (num * xj) % prime实际算出的拉格朗日系数与标准公式差一个符号,是一个隐蔽bug。
修正后的重构循环如下:
# 修正:分子必须包含符号 num = (num * (-xj)) % prime对应的完整reconstruct_secret在本地运行时会作为整个代码库的基础函数。测试时用随机秘密和阈值组合循环几百次,每次都应该恢复出原始秘密。
3.3 浮点梯度与有限域整数的编解码
# quant.py SCALE = 10000 # 量化系数,可调节 PRIME = 2**61 - 1 # 梅森素数,模运算快且足够大 def float_to_field(vec, scale=SCALE, prime=PRIME): """浮点向量转有限域元素,负数通过取模并入 [0, prime)""" return [int(round(v * scale)) % prime for v in vec] def field_to_float(vec, scale=SCALE, prime=PRIME): """有限域元素还原为浮点,超过 prime/2 视为负数""" out = [] for v in vec: v = v if v < prime // 2 else v - prime out.append(v / scale) return out量化这一层是FedSTSS精度损失的唯一天然来源。scale取10000,意味着梯度被保留到小数点后四位,对于大多数联邦学习模型足够;但如果你的模型参数范围本身就很大,比如某些归一化前的特征权重绝对值到100以上,就必须调小scale或者改用分段量化。另外要注意int(round(v * scale))%prime对负数取模的结果会自动落到[0, prime),编码阶段不需要手动处理负号。
3.4 客户端侧:掩码生成与份额交换
客户端侧的逻辑可以抽象为一个函数:输入本地更新向量和从其他客户端收到的份额列表,返回带掩码的更新和一条汇总份额向量。
# client_side.py import random from shamir import generate_shares def client_side(delta_w, self_id, threshold, num_clients, prime, received_shares, scale=SCALE): # 1. 量化并生成掩码 delta_int = [int(round(v * scale)) % prime for v in delta_w] mask = [random.randrange(prime) for _ in range(len(delta_int))] masked = [(d + m) % prime for d, m in zip(delta_int, mask)] # 2. 对掩码向量的每个元素做 Shamir 分割 # my_shares[k][x] 表示第 k 个掩码值分给参与方 x 的份额 my_shares = [] for m in mask: my_shares.append(generate_shares(m, threshold, num_clients, prime)) # 3. 把自己的 self_id 份也加入汇总,再叠加其他客户端发来的份额 # received_shares 的元素是其他客户端计算出的 my_shares bonus = [] for k in range(len(mask)): total = 0 for sender_shares in received_shares: total = (total + sender_shares[k][self_id][1]) % prime bonus.append(total) return masked, bonus这里的received_shares是通信层组装好以后传给函数的数据结构。每收到一个客户端的全部份额,就按k取出第k个掩码值的第self_id份。因为所有客户端的份额x坐标相同,逐项相加后送到服务器,服务器才能用插值恢复掩码总和。注意客户端自己也应该把自己的份额加进去,这里实现为received_shares里包含自身份额。
3.5 服务器侧:恢复掩码总和并完成平均
# server_side.py from shamir import reconstruct_secret def server_aggregate(masked_updates, bonus_vectors, online_ids, threshold, prime, scale=SCALE): d = len(masked_updates[0]) # 1. 用在线客户端的 bonus 向量做拉格朗日重构 mask_sum = [] for k in range(d): points = [(online_ids[i] + 1, bonus_vectors[i][k]) for i in range(threshold)] mask_sum.append(reconstruct_secret(points, prime)) # 2. 带掩码更新直接相加 total_masked = [0] * d for vec in masked_updates: for i, v in enumerate(vec): total_masked[i] = (total_masked[i] + v) % prime # 3. 解掩码并平均 plain_sum = [(a - b) % prime for a, b in zip(total_masked, mask_sum)] avg = field_to_float(plain_sum, prime, scale) return [x / len(masked_updates) for x in avg]server_aggregate里的阈值即t,online_ids是实际在线并提供有效份额的客户端编号列表。这里取前threshold个参与点做插值,如果某个客户端掉线导致可用点不足t个,本轮聚合会直接失败,这也是t参数和容错能力的核心约束。完整工程还需要在客户端训练部分用numpy实现逻辑回归或小规模MLP的梯度计算,这与普通FedAvg的客户端训练完全一致,区别只在提交前的掩码与份额处理。
4. 对比实验设计:FedSTSS vs FedAvg与Paillier方案的精度与开销
4.1 实验配置与数据集选择
对比实验的目标是回答三个问题:安全聚合方案相对FedAvg牺牲了多少精度;通信和计算开销在什么量级;门限参数对结果的影响是否显著。数据集我会用scikit-learn自带的digits手写数字集,样本量约1800个,特征维度64,完全可以支撑一轮快速验证。模型用逻辑回归,梯度用numpy手写,避免引入深度学习框架后掩盖安全聚合本身的耗时。
参与方数量设为8,每轮随机选取其中6个参与训练,模拟真实联邦学习的部分参与场景。FedSTSS的门限t设置为5,意思是允许本轮最多1个客户端掉线且不泄露掩码。所有方案使用相同的全局模型初始化,随机种子固定,否则不同方案之间的精度差异会混入初始化噪声。代码入口如下:
python run_comparison.py --scheme fedavg --clients 8 --rounds 50 python run_comparison.py --scheme fedstss --clients 8 --threshold 5 --rounds 50 python run_comparison.py --scheme paillier --clients 8 --rounds 10Paillier方案只跑10轮,是因为phe库对每个整数参数执行加密的耗时在毫秒级以上,64维特征乘以10个类别就有640个参数,8个客户端完整跑50轮会非常慢。这本身就是同态加密方案在性能上的一个重要观察点。
4.2 三种方案的通信与计算差异对比
| 对比项 | FedAvg | FedSTSS(t=5, n=8) | Paillier加密聚合 |
|---|---|---|---|
| 单客户端上传量 | 一次性上传明文更新 | 带掩码更新 + 一条汇总份额 | 每个参数一个密文 |
| 服务器聚合操作 | 明文平均 | 拉格朗日插值 + 模加 | 密文加法 |
| 是否泄露单客户端梯度 | 是 | 否 | 否 |
| 掉线容忍 | 任意数量 | 最多n-t个 | 任一掉线则失败 |
| 精度损失来源 | 无 | 量化误差 | 无 |
通信量的量级可以参考:假设模型参数d个,参与方n个。FedAvg上传n·d个浮点数;FedSTSS上传n·d个带掩码整数再加n·d个汇总份额,总量约2倍;Paillier则是n·d乘上密文扩张系数,通常每个32位整数的密文要占256字节以上。在digits这种小模型上差异还不明显,换到千万元素的大模型时,通信和计算差距会被放大到完全不可接受的程度。
4.3 对比实验脚本的骨架
# run_comparison.py 关键片段 import argparse import numpy as np from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split def get_clients(X, y, n_clients): """把数据集均匀切分给 n_clients 个客户端""" idx = np.arange(len(X)) np.random.shuffle(idx) return [idx[i::n_clients] for i in range(n_clients)] def evaluate(global_w, X, y): """逻辑回归准确率""" pred = np.argmax(X @ global_w, axis=1) return np.mean(pred == y) def run_fedstss(clients, global_w, X_test, y_test, threshold, rounds): for r in range(rounds): masked_updates, bonus_vectors, online_ids = [], [], [] for cid in clients: # 每轮随机挑选在线客户端 ... delta = local_update(...) masked, bonus = client_side(delta, cid, threshold, len(clients), PRIME, received_shares) masked_updates.append(masked) bonus_vectors.append(bonus) avg = server_aggregate(masked_updates, bonus_vectors, online_ids, threshold, PRIME) global_w = global_w - lr * np.array(avg)脚本里local_update需要自己实现一个批梯度下降函数,输入客户端样本和全局权重,返回梯度。为了公平,FedAvg和FedSTSS必须共用同一个local_update和相同的学习率,唯一差别只在提交更新前是否做掩码和份额交换。Paillier方案在聚合前把delta_int逐个加密,服务器端对密文做加法后由可信第三方解密。需要注意Paillier不支持减法,负梯度需要预先映射到非负整数区间,否则解密结果会出错。
4.4 实验结果记录与量化误差排查
实验记录推荐固定三张表:精度随轮次变化、单轮平均通信字节数、单轮聚合耗时。通信字节数可以用pickle.dumps后取len来统计,耗时用time.perf_counter扣掉客户端本地训练时间单独测量。如果FedSTSS的最终精度与FedAvg差超过0.5个百分点,优先怀疑是scale取值太小导致量化误差过大;把scale从10000调整到100000再跑一轮,通常精度就能对齐。如果精度反而震荡,则检查field_to_float的负数还原逻辑是否在聚合值超过p/2时被误判。
5. 门限t与参与方n的平衡:容错边界与验证技巧
5.1 t值的选择策略与场景映射
门限t决定了联邦系统的安全与可用边界。从安全角度看,攻击者至少需要拿到t份关于同一掩码的份额才能恢复该掩码;从可用角度看,本轮在线客户端必须大于等于t,否则聚合无法完成。设最大允许掉线数为f,同时希望保持门限安全性,则t必须同时满足t ≤ n - f和t > f,也就是f < t ≤ n - f。实际工程里常见的几档配置如下:
| 门限t | 可容忍掉线数 | 可容忍泄露份额数 | 适用场景 |
|---|---|---|---|
| t = n | 0 | n-1 | 强保密,参与方全部在线 |
| t = n - 1 | 1 | n-2 | 高可靠核心节点场景 |
| t = n/2 + 1 | n/2 - 1 | n/2 | 兼顾容错与安全的默认选择 |
| t = 2 | n-2 | 1 | 小规模合作,信任要求高 |
我一般会优先取t = n/2 + 1,这样在n个参与方里掉线不超过一半时训练都能继续,同时任何少于半数参与方的合谋也拿不到掩码。如果你的场景是两家机构对等合作,t = 2即可,但要注意此时只要一个参与方泄露份额,掩码就有可能被恢复。
5.2 用脚本验证重构正确性与加法同态
Shamir实现的最大隐患是拉格朗日插值在边界条件下出错,尤其是份额点顺序打乱、x坐标不从1开始、或临时加入了负数坐标。写一个轮询测试函数把常见情况覆盖住:
# verify.py import random from shamir import generate_shares, reconstruct_secret def roundtrip_test(threshold, num_clients, prime, repeat=500): for _ in range(repeat): secret = random.randrange(1, prime) shares = generate_shares(secret, threshold, num_clients, prime) assert reconstruct_secret(shares[:threshold], prime) == secret # 随机取 threshold 份,而不是固定前 threshold 份 picked = random.sample(shares, threshold) assert reconstruct_secret(picked, prime) == secret print("roundtrip ok") def homomorphic_test(threshold, prime): m1, m2 = random.randrange(prime), random.randrange(prime) s1 = generate_shares(m1, threshold, threshold, prime) s2 = generate_shares(m2, threshold, threshold, prime) points = [(x, (y1 + y2) % prime) for (x, y1), (_, y2) in zip(s1, s2)] assert reconstruct_secret(points, prime) == (m1 + m2) % prime print("homomorphic ok")roundtrip_test验证基础的秘密恢复,homomorphic_test验证FedSTSS真正依赖的加法同态性质,即两个掩码份额逐项相加后再插值,结果等于两个掩码之和。这两组测试我建议在跑任何实验前先执行,因为它们能同时排除掉实现的符号错误和坐标错位问题。
5.3 量化参数的自检技巧
最后一个容易在项目验收时被问到的点是scale与prime的配合。把上面代码跑通后,可以加一段溢出检查:在所有客户端更新编码后,随机挑一个维度,计算所有掩码与更新绝对值的和,确认其低于prime的一半。一旦累加和超过prime的1/4,就应该增大prime或降低scale。这一步虽然简单,但能避免最隐蔽的模回绕错误——精度看起来只差一点,实际聚合结果已经完全错乱。
本文还有配套的精品资源,点击获取