
简介面向安全聚合与联邦学习研究场景资源提供基于Shamir门限秘密共享的联邦学习安全聚合模型FedSTSS的完整Python实现包含与FedAvg、FedShare等主流方法的对比实验适用于计算机、人工智能相关专业的毕设、课设或科研复现。资源包共54个文件以22个Python脚本为核心覆盖联邦平均/共享/门限秘密共享的客户端-服务器实现、Shamir算法模块、通用数据与模型工具另有6个Shell启动脚本、20个日志记录、4份CSV数据集及README文档方便快速运行与核对实验过程。压缩包仅68KB代码轻量易读目前已有140人学习浏览。通过源码可清晰理解秘密共享如何在联邦学习聚合中实现隐私保护配合日志与数据集可复现对比结果基础扎实者还能基于现有模块扩展新机制或数据集用于论文实验或项目演示。1. FedSTSS 是什么安全聚合最先要回答凭什么信任服务端联邦学习Federated Learning的标准流程是客户端在本地训练模型只把梯度或模型参数上传给服务端服务端做平均后再下发。这个设计默认服务端是诚实的但现实中一个能读到明文梯度的服务端完全可以通过梯度反推训练数据甚至构造恶意更新污染全局模型。基于Shamir门限秘密共享的联邦学习安全聚合模型FedSTSS走的是另一条路把每个客户端的梯度先拆成若干份额份额分散交给多个参与方只有收集到足够份额才能恢复出梯度——而这个恢复恰好就等价于一次安全的聚合。换句话说FedSTSS让隐私保护不依赖对服务端的信任而是依赖数学上的门限结构。适合谁正在做联邦学习隐私保护的Python研究者需要提交联邦学习代码作为课设或竞赛方案的在校生以及在企业里想把聚合服务做成无法偷看明文的那类工程师。这篇文章从原理讲到代码实现再讲到对比实验怎么设计、参数怎么设、坑在哪照着这条路径能把FedSTSS从一个名字变成你能跑、能改、能答辩的项目。2. Shamir门限秘密共享的原理与选型为什么偏偏是拉格朗日插值2.1 从多项式到份额t-out-of-n 到底怎么运作Shamir门限秘密共享的核心思想说穿了就是一条中学数学结论任意t个点可以唯一确定一条t-1次多项式曲线。要把一个秘密拆成n份构造一个t-1次多项式常数项等于秘密值其余t-1个系数随机取然后把不同的x坐标代入多项式得到n个点。这n个点就是n份份额分给n个参与方每人一个。想恢复秘密只需要拿到任意t个点做拉格朗日插值常数项就被精确算回来。少于t个点时能拟合出无穷多条多项式秘密在信息论意义上完全不可知。在联邦学习里秘密不是一个文件而是某个客户端本地的梯度向量。整段梯度没办法直接塞进一个多项式常数项所以常见做法是把梯度展平成一维数组对每一维分别做Shamir份额化另一种做法是按神经网络层分块对每一块的参数张量做份额化。一个客户端把自己的梯度分给其它客户端和服务端各一份自己保留一份。在FedSTSS的安全聚合中每个参与方拿到的是别人的梯度份额不是别人的明文梯度。服务端做完插值恢复恢复出来的结果已经是多个客户端份额的聚合值单个客户端的明文梯度在这个过程中从来没有完整出现过。提示这里有一个很关键的理解点——份额可以在参与方之间做本地聚合。也就是说每个参与方先把收到的所有份额相加再上传这个聚合后的份额服务端最后只需要一个插值步骤就能得到全局聚合值而不是为每个客户端单独恢复梯度再平均。这个设计让通信量和计算量都大幅下降。2.2 为什么不用同态加密和差分隐私做安全聚合有三条主流技术路线理解FedSTSS为什么选Shamir得先看另外两条为什么不选。同态加密可以把密文相加得到加密的和服务端在不解密的情况下完成聚合。但代价非常现实计算开销和通信量都大到让人退缩。客户端几百维的梯度经过同态加密后体积放大几十倍甚至上百倍每轮通信上传的都是密文训练一轮的时间能拖到普通联邦学习的几十倍。差分隐私则是相反的方向在梯度上添加噪声再上传计算开销小但加噪会直接伤害模型精度收敛变慢甚至不收敛噪声预算的分配也是个玄学。Shamir门限方案的计算开销主要在多客户端之间分发份额的通信上服务端聚合本身是解一个插值方程不涉及重量级密码运算Python里的int类型天然支持大整数模运算实现成本低。它也有自己的代价需要参与方之间额外通信掉线数量超过n-t就会恢复失败而且必须假设做份额中继的参与方不会相互勾结。对联邦学习这种客户端本来就是分布式、天然有通信链路的场景这些代价是能接受的。2.3 素数域与模数选择最容易把实现带进坑的位置Shamir份额的加减乘除必须在一个素数域里完成否则插值恢复会混入浮点误差跑两轮之后模型参数全乱。 常用做法是选一个大素数作为模数比如形如2^k-1的梅森素数Python里直接用pow(2, 127) - 1就能拿到。 所有除法都用模逆元实现。为什么必须是素数因为模逆元存在的前提是分母与模数互质而素数域里任意非零元素都有逆元。选了合数做模数某些份额在求逆时会直接失败聚合当场翻车。模数的位数直接决定安全性。128位模数在目前的安全强度下足够没有必要用512位。模数越大份额值越大通信数据量也跟着涨。我在实现里默认用2^61-1原因是它比Python int的常规小整数运算边界小几个数量级调试时肉眼能看出错误上传时也能压缩得更小。想更稳就切到2^127-1代码只需要改一个常量其它逻辑不用动。2.4 门限参数选择t和n到底怎么定n等于参与聚合的参与方总数通常包括全部客户端加服务端t是恢复秘密所需的最少份额数。安全性和可用性在此处互相拉扯t越大越安全因为恶意参与者要凑够t个份额才能勾结出秘密t越小越抗掉线因为掉线客户端不超过n-t就能恢复。常见做法是把t设为n的2/3到3/4。例如10个客户端加服务端共11个参与方t取7或8比较均衡。7意味着要勾结7个参与方才能解开秘密同时允许4个参与方掉线不影响训练。如果t取到n安全上倒是拉到最满但任何一方掉线全局停摆这在真实网络环境下没法用。如果t取太小比如t3那三个参与方联合起来就能拿到聚合梯度隐私保护形同虚设。做对比实验时我一般会把t本身也作为一个变量测观察它对精度和鲁棒性的影响。3. 跑通 FedSTSS环境准备、数据集划分与最小复现路径3.1 环境依赖与安装三步装好运行环境FedSTSS的完整实现依赖Python生态里最常见的几个库没有重型的密码学依赖。Python版本建议3.9以上PyTorch负责模型训练NumPy负责张量操作和份额生成scikit-learn用来做数据集切分和评估指标计算。# 创建虚拟环境避免污染系统Python python -m venv fedstss_env source fedstss_env/bin/activate # 安装核心依赖按顺序执行 pip install torch2.1.0 numpy1.26.2 scikit-learn1.3.2第一行创建独立的虚拟环境这是Python项目的标准做法避免多项目间的依赖版本冲突。安装时固定了三个主库版本因为新版PyTorch对旧版NumPy的兼容性时有波动固定版本能保证下面所有代码原样可跑。如果你的机器装的是CPU版PyTorch训练完全够用FedSTSS的瓶颈本来就在份额分发和聚合不在单机训练速度上。3.2 项目结构一份按模块拆好的代码骨架拿到FedSTSS源码后先别急着运行花两分钟看懂目录结构。这个项目的文件组织是典型的算法实现与实验分离结构调参和换数据集都只动部分文件。fedstss_project/ ├── config.py # 全局参数客户端数、门限值、模数、学习率 ├── shamir_secret.py # Shamir份额生成与拉格朗日恢复核心模块 ├── fedstss_aggregator.py # 服务端聚合逻辑调用shamir恢复 ├── client.py # 客户端本地训练与梯度份额化 ├── dataset_partition.py # 数据切分模拟非独立同分布 ├── train_fedstss.py # 主训练入口FedSTSS完整流程 ├── train_fedavg.py # 对照组明文FedAvg ├── experiments.py # 对比实验脚本输出精度与误差统计 └── results/ # 实验结果输出目录shamir_secret.py是整套代码的核心份额生成和恢复都在这里任何改动会直接影响聚合正确性。config.py集中了所有可调参数门限值、客户端数、模数、训练轮次都改这一个文件。train_fedavg.py是必须的对照组没有它对比实验就没有基线后面评判FedSTSS的精度损失也无从谈起。3.3 数据准备把IID数据切出非独立同分布的样貌联邦学习的数据分布和普通深度学习不一样。普通训练把数据随机打乱切train/test联邦学习要模拟每个客户端的数据分布不同这个现实最常见的就是非独立同分布Non-IID切分。比如医院A的数据全是某种疾病的影像医院B另一种疾病各自数据分布完全不同。FedSTSS实验里用的是最经典的按标签分布切分法。import numpy as np from sklearn.datasets import fetch_openml # 以MNIST为例把10个类别的数据按不同分布分给10个客户端 def partition_non_iid(labels, num_clients10, num_shards200, num_classes10): # 第一步按类别把样本索引分组 class_indices [np.where(labels i)[0] for i in range(num_classes)] # 第二步每个类别均分成num_shards/num_classes个分片 shards [] for class_idx in class_indices: shards.extend(np.array_split(class_idx, num_shards // num_classes)) # 第三步把分片随机分配给客户端每个客户端拿num_shards/num_clients个分片 np.random.shuffle(shards) client_data {} for client_id in range(num_clients): assigned np.concatenate( shards[client_id * (num_shards // num_clients): (client_id 1) * (num_shards // num_clients)] ) client_data[client_id] assigned return client_data if __name__ __main__: from sklearn.datasets import fetch_openml mnist fetch_openml(mnist_784, version1, as_frameFalse, parserauto) labels mnist.target.astype(int) distribution partition_non_iid(labels) print(f客户端数量: {len(distribution)}, 每个客户端样本量约为: {len(distribution[0])})这段代码的核心逻辑是把MNIST的10个类别切分成200个分片然后随机分给10个客户端。每个客户端会拿到属于不同类别的分片类别分布差异明显这就模拟出了非独立同分布的效果。num_shards越大每个客户端的类别越杂数据分布越接近独立同分布num_shards越小每个客户端的数据类别越单一。200这个数值是经验值10个客户端、10个类别时每个客户端拿20个分片大约覆盖2个主要类别分布倾斜明显但又不至于让某个客户端完全只有一类数据导致训练不收敛。4. 核心模块实现份额生成、分发聚合与关键参数设定4.1 Shamir份额生成与恢复手写最小可用实现先看shamir_secret.py里的核心函数。这段代码是整个项目的数学地基后面的聚合逻辑全部建立在这上面。import random from typing import List, Tuple DEFAULT_MODULUS (1 127) - 1 # 2^127 - 1梅森素数作为素数域模数 def split_secret(secret: int, threshold: int, num_shares: int) - List[Tuple[int, int]]: 把整数秘密拆成num_shares份threshold份可恢复。 if secret DEFAULT_MODULUS: raise ValueError(f秘密值必须小于模数当前模数为 {DEFAULT_MODULUS}) if threshold num_shares: raise ValueError(f门限值 {threshold} 不能超过份额数 {num_shares}) coefficients [secret] # 常数项即秘密本身 # 随机生成 threshold-1 个系数保证多项式t-1次 coefficients.extend(random.randint(0, DEFAULT_MODULUS - 1) for _ in range(threshold - 1)) shares [] for x in range(1, num_shares 1): # x从1开始避免x0泄露秘密 # 用霍纳法求多项式值避免大数乘方性能爆炸 y coefficients[-1] for coeff in reversed(coefficients[:-1]): y (y * x coeff) % DEFAULT_MODULUS shares.append((x, y)) return shares def recover_secret(shares: List[Tuple[int, int]], threshold: int) - int: 用拉格朗日插值恢复秘密即多项式在x0处的值。 if len(shares) threshold: raise ValueError(f份额数 {len(shares)} 不足门限 {threshold}无法恢复) secret 0 for i, (xi, yi) in enumerate(shares): numerator 1 denominator 1 for j, (xj, _) in enumerate(shares): if i j: continue # 拉格朗日基多项式目标点是x0 numerator (numerator * (-xj)) % DEFAULT_MODULUS denominator (denominator * (xi - xj)) % DEFAULT_MODULUS # 模逆元实现除法这是素数域运算的关键 lagrange_coeff (numerator * pow(denominator, -1, DEFAULT_MODULUS)) % DEFAULT_MODULUS secret (secret yi * lagrange_coeff) % DEFAULT_MODULUS return secretsplit_secret构造多项式时常数项直接设为秘密值其余系数随机生成。份额的x坐标从1开始因为x0会直接暴露常数项。恢复时用拉格朗日基多项式计算x0处的函数值所有的除法和加法都在模数下进行。pow(denominator, -1, DEFAULT_MODULUS)是Python 3.8之后的特性一行代码完成模逆元计算比手写扩展欧几里得算法简洁得多而且不会出错。为什么要用霍纳法而不是sum(coeff * x**i)因为大整数乘方的计算量随指数增长在模数127位时性能差距明显霍纳法把乘方次数压到与多项式次数相等循环里每次只做一次乘法和一次加法。4.2 安全聚合全流程客户端的份额化与服务端的插值恢复有了份额生成函数FedSTSS的完整聚合流程分四步走。第一步每个客户端在本地完成一轮训练拿到梯度向量第二步对梯度向量的每一维做Shamir份额化得到n份子份额第三步参与方之间交换份额各方把收到的所有份额按维度相加得到一个聚合份额第四步上传聚合份额给服务端服务端收集到t个聚合份额后做一次拉格朗日插值恢复出全局梯度。import numpy as np from shamir_secret import split_secret, recover_secret, DEFAULT_MODULUS class FedSTSSClient: 客户端角色本地训练 梯度份额化 def __init__(self, client_id: int, num_participants: int, threshold: int): self.client_id client_id self.num_participants num_participants self.threshold threshold def share_gradient(self, grad_vector: np.ndarray) - dict: 把梯度向量逐维拆成份额返回给其它参与方的份额映射。 shares_by_participant {i: [] for i in range(self.num_participants)} for grad_value in grad_vector: # 浮点梯度转整数先缩放再取整保留6位小数精度 secret_int int(round(float(grad_value) * 1e6)) % DEFAULT_MODULUS shares split_secret(secret_int, self.threshold, self.num_participants) for idx, (part_id, share_value) in enumerate(shares): shares_by_participant[idx].append(share_value) return shares_by_participant def aggregate_local_shares(self, received: list) - np.ndarray: 把收到的所有份额按维度相加得到聚合份额。 received_array np.array(received, dtypeobject) # object类型防止大整数溢出 return received_array.sum(axis0) % DEFAULT_MODULUS这里有两个细节必须注意。第一梯度是浮点数Shamir份额化只能处理整数所以先用缩放取整把浮点转成整数。缩放系数1e6意味着我们保留6位小数精度。要保留更多精度就放大倍数但代价是秘密值变大逼近模数上限的风险增加。第二dtypeobject非常重要NumPy默认用int64存储而份额值在127位模数下远超int64上限不指定object类型结果会被静默截断聚合结果完全错误。服务端聚合的逻辑更简单只需要收集聚合份额并调用恢复函数class FedSTSSAggregator: 服务端角色收集聚合份额恢复全局梯度 def __init__(self, threshold: int): self.threshold threshold def aggregate(self, share_lists: dict) - np.ndarray: share_lists: {参与方id: [该参与方收到的一维聚合份额]} dims len(next(iter(share_lists.values()))) # 服务端只需要任意threshold个参与方上传的聚合份额 collected [] for participant_id, shares in share_lists.items(): shares_tuple [(participant_id, s) for s in shares] collected.append(shares_tuple) if len(collected) self.threshold: break recovered_dims [] for dim in range(dims): # 每个维度取threshold个参与方的份额做恢复 selected [shares[dim] for shares in collected] secret_int recover_secret(selected, self.threshold) # 整数转回浮点梯度除回缩放系数 recovered_dims.append(secret_int / 1e6 if secret_int DEFAULT_MODULUS / 2 else (secret_int - DEFAULT_MODULUS) / 1e6) return np.array(recovered_dims, dtypefloat)最后的符号处理是关键步骤。模数域里的大整数可能是负数映射过来的如果秘密值超过模数的一半转换回浮点时要减掉模数恢复成负值。这个判断在普通课程项目里常常被漏掉结果就是聚合出来的梯度全变成大正数模型直接发散。4.3 核心参数表一份可以直接抄的配置清单下面这张表总结了FedSTSS实现里最核心的参数和默认值每项都标注了调整方向和风险。参数默认值作用调参注意num_clients10参与训练的客户端总数越大通信开销越高门限也要跟着调threshold (t)7恢复秘密所需最少份额数偏大安全但脆弱偏小鲁棒但易勾结modulus2^127-1素数域模数所有运算的边界改小会溢出改大徒增通信量scale_factor1e6浮点梯度转整数的缩放系数过大逼近模数上限过小丢精度local_epochs1每轮本地训练次数联邦学习默认一轮多轮加剧分布偏移learning_rate0.01SGD学习率FedSTSS不改变训练本身按普通设定最值得花时间调的是threshold因为它直接决定安全性和可用性的平衡。建议做一组threshold从5到9的对比实验观察最终精度和掉线容忍度的关系这也会是你答辩时最有说服力的数据。5. 避坑指南Shamir安全聚合最常见的五个翻车点5.1 现象恢复出来的梯度全是负数大数与明文聚合结果完全对不上第一轮训练结束服务端恢复出全局梯度更新模型后loss不降反升打印梯度一看全是几亿量级的大整数。原因几乎可以锁定在符号处理上模数域里的大整数表示的是负数但恢复代码没有做超过模数一半就减模数的反映射。解决方法是按照4.2节里的写法在整数转回浮点之前判断秘密值是否大于模数的一半是则减去模数。这个坑我踩过一次之后现在每次写完恢复函数第一件事就是构造一个已知的负梯度做单元测试。5.2 现象明明份额数是够的拉格朗日恢复出来的秘密却不对训练过程一切正常突然某轮聚合结果不对但重跑一次又对了。这是随机性问题Shamir份额生成依赖随机系数而拉格朗日插值对份额点的质量敏感度为零——只要份额点都在多项式上任意t个都能恢复成功。真正的问题通常出在不同客户端用了不同的随机数状态导致同一轮训练里多项式构造不一致。解决方法是给每个客户端设置固定的随机种子种子由config.py统一分发同时保证所有客户端在训练开始时同步随机数状态。这个修复能消除大部分时好时坏的诡异现象。5.3 现象客户端掉线一多服务端直接抛异常崩溃真实网络环境下掉线是常态但代码如果一遇到掉线就抛ValueError整个训练流程就断了。原因是聚合逻辑没有做掉线容忍设计。解决方法是把客户端掉线当成正常事件处理服务端在等待聚合份额时设置超时收集到threshold份就开始聚合忽略其余未到达的参与方。掉线数量超过n-t时无法恢复秘密这时可以跳过本轮聚合沿用上一轮全局模型继续训练下一轮。在代码层面要给聚合函数加一个timeout参数超时后检查当前收集的份额数达到门限才插值否则返回上个模型权重。5.4 现象门限设成了n所有客户端必须在线才能训练有一种误区是门限越大越安全于是把threshold设成和num_participants相等。这种做法在单机模拟里一切正常因为所有进程都在本机跑不存在真正的掉线。一旦部署到多机或模拟真实掉线场景任何一个客户端掉线都会导致整轮训练失败系统可用性变成零。解决方法是把threshold默认设为参与方总数的2/3到3/4并且在对比实验里展示不同threshold下的训练曲线证明这个选择是权衡后的结果不是随手写的。5.5 现象用浮点运算做拉格朗日插值梯度精度灾难性丢失有同学图省事觉得模逆元麻烦直接用浮点除法做拉格朗日插值。结果在小数据集上还能跑数据集一换模型完全训不动。原因是浮点插值在高次多项式场景下会引入不可控的舍入误差t越大误差越大。解决方法是所有份额运算严格走素数域除法用模逆元。浮点只出现在两个边界梯度转整数之前、恢复整数转回浮点之后。中间过程不允许任何float参与。6. 验证聚合正确性的两个小技巧与压测习惯FedSTSS这类安全聚合实现比普通联邦学习多做了一层数学变换所以验证不能只看最终精度还要验证数学变换本身没有引入错误。我每次拿到新代码先不跑完整训练而是构造一个5维的小梯度向量手动执行份额化→聚合份额→插值恢复三个阶段断言恢复结果和原始梯度的误差小于1e-5。这个测试5分钟就能跑完能挡掉上面提到的符号问题、模数溢出问题、缩放出错问题。第二个技巧是跑一个明文对账实验。同一种数据分布和初始化条件下分别用FedAvg和FedSTSS跑完整训练每轮对比两者聚合出来的梯度余弦相似度。正常情况下梯度应该高度相似只是精度有微小差异。如果某轮相似度突然掉到0.9以下说明这个batch的份额交换或聚合公式有bug可以再缩小范围定位。这个方法比只看最终accuracy敏感得多能在第3轮就发现错误而不是等50轮训练完才发现整个实验报废。压测方面建议做三件小事一是模拟随机掉线每轮随机drop掉1到2个客户端确认精度曲线不会发散二是把threshold从小往大扫描确认达到门限才能恢复这个边界行为三是多跑几个随机种子确认结果可复现。这三件小事加起来不到半小时但会让你的对比实验结论从可能巧合变成统计显著。我做这个项目时最大的教训就是太早跑完整训练翻车之后返工调试的成本远高于一开始写好那几个小测试。先证明数学层是对的再训练模型这个顺序不能颠倒。希望这些经验和踩坑记录能帮你在FedSTSS上少走几趟弯路。本文还有配套的精品资源点击获取