
简介面向联邦学习安全聚合研究的一套可运行代码实现重点给出基于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, scaleSCALE, primePRIME): 浮点向量转有限域元素负数通过取模并入 [0, prime) return [int(round(v * scale)) % prime for v in vec] def field_to_float(vec, scaleSCALE, primePRIME): 有限域元素还原为浮点超过 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, scaleSCALE): # 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, scaleSCALE): 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里的阈值即tonline_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 三种方案的通信与计算差异对比对比项FedAvgFedSTSS(t5, n8)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, axis1) 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 n0n-1强保密参与方全部在线t n - 11n-2高可靠核心节点场景t n/2 1n/2 - 1n/2兼顾容错与安全的默认选择t 2n-21小规模合作信任要求高我一般会优先取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, repeat500): 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。这一步虽然简单但能避免最隐蔽的模回绕错误——精度看起来只差一点实际聚合结果已经完全错乱。本文还有配套的精品资源点击获取