联邦学习中的边缘模型聚合:安全聚合协议与差分隐私的 Rust 实现

📅 2026/7/22 10:20:52
联邦学习中的边缘模型聚合:安全聚合协议与差分隐私的 Rust 实现
联邦学习中的边缘模型聚合安全聚合协议与差分隐私的 Rust 实现一、模型参数明文传输的安全风险联邦学习的标准流程N 个边缘设备各自在本地数据上训练模型将梯度更新上传到中央服务器进行聚合FedAvg服务器平均后下发新模型。问题在于梯度是明文上传的。攻击者可以执行梯度泄露攻击——通过分析梯度重建训练数据。一项知名的攻击演示给定一个 batch 的梯度更新攻击者可以重建 batch 中的原始图像Deep Leakage from Gradients, 2019。对于医疗、金融等隐私敏感行业这意味着联邦学习名义上数据不出设备实际上等价于明文共享训练数据。两项技术组合解决此问题安全聚合Secure Aggregation通过多方安全计算MPC协议服务器只能得到梯度的聚合值求和看不到单个设备的梯度。差分隐私Differential Privacy在梯度中注入精心设计的噪声即使攻击者拿到了聚合后的梯度也无法推断任何单个训练样本的信息。安全聚合通过密钥协商和秘密共享实现服务器只能看到聚合结果。差分隐私通过高斯噪声扰动实现聚合结果不泄露个体信息。两者正交互补——安全聚合保护传输过程差分隐私保护聚合结果。二、安全聚合与差分隐私的协同架构协议分四个阶段密钥协商Key Agreement每对设备 (i, j) 通过 Diffie-Hellman 协议协商一个共享密钥s_{i,j}。基于此密钥生成 pairwise maskprg(s_{i,j})。Masking掩码设备 i 的真实梯度ΔW_i加上与所有其他设备的 pairwise mask设备 i 加prg(s_{i,j})设备 j 减prg(s_{i,j})。这样 pairwise mask 在聚合时相互抵消因为 i 加了 j 的 maskj 减了 i 的 mask服务器只能看到真实梯度的总和。Unmasking解掩码当设备掉线时这在移动边缘设备中很常见其配对 mask 无法抵消。此时存活的设备需要上传掉线设备的秘密份额供服务器重建 mask 并去除。聚合与加噪服务器聚合解密后的梯度并施加差分隐私机制。高斯噪声的标准差由 privacy budget ε 决定——ε 越小隐私保护越强但模型精度下降越多。三、安全聚合与差分隐私的 Rust 实现use rand::RngCore; use rand::rngs::OsRng; use sha2::{Sha256, Digest}; use hkdf::Hkdf; use std::collections::HashMap; use std::sync::Arc; use tokio::sync::Mutex; /// 设备 ID type DeviceId u64; /// 梯度向量 —— 每层展平为一维 type Gradient Vecf32; /// 安全聚合客户端运行在边缘设备上 pub struct SecureAggClient { /// 当前设备 ID device_id: DeviceId, /// 与该设备配对的密钥: key 对方 device_id, value 共享密钥 /// HKDF 派生s_{i,j} HKDF(DH(sk_i, pk_j)) pairwise_keys: HashMapDeviceId, Vecu8, /// 本设备的私钥 —— ECDH P-256 private_key: p256::SecretKey, /// 差分隐私参数 dp_config: DpConfig, } /// 差分隐私配置 pub struct DpConfig { /// 隐私预算 ε —— 越小噪声越大 pub epsilon: f64, /// δ 参数 —— (ε, δ)-DP 的松弛项 pub delta: f64, /// 梯度裁剪范数 C —— 限制单样本梯度的最大范数 pub clip_norm: f32, } /// 安全聚合服务端运行在中央服务器上 pub struct SecureAggServer { /// 所有存活设备的公钥 public_keys: HashMapDeviceId, p256::PublicKey, /// 聚合后的梯度 aggregated: MutexOptionGradient, /// 存活设备列表 alive_devices: MutexVecDeviceId, } // 1. 密钥协商 impl SecureAggClient { /// 初始化 —— 生成密钥对 pub fn new(device_id: DeviceId, dp_config: DpConfig) - Self { let private_key p256::SecretKey::random(mut OsRng); Self { device_id, pairwise_keys: HashMap::new(), private_key, dp_config, } } /// 获取公钥 —— 在密钥协商阶段发送给服务器 pub fn public_key(self) - p256::PublicKey { self.private_key.public_key() } /// 建立与另一设备的共享密钥 /// 使用 ECDH: s DH(my_sk, peer_pk) /// 再通过 HKDF 派生为 AES 密钥 pub fn establish_pairwise_key( mut self, peer_id: DeviceId, peer_pk: p256::PublicKey, ) { // ECDH 密钥交换 let shared_secret p256::ecdh::diffie_hellman( self.private_key.to_nonzero_scalar(), peer_pk.as_affine(), ); // HKDF 派生 —— 将原始 DH 共享密钥转换为固定长度 AES 密钥 // 使用 peer_id device_id 作为 salt确保方向的确定性 let salt if peer_id self.device_id { [peer_id.to_le_bytes(), self.device_id.to_le_bytes()].concat() } else { [self.device_id.to_le_bytes(), peer_id.to_le_bytes()].concat() }; let hkdf Hkdf::Sha256::new(Some(salt), shared_secret.as_bytes()); let mut okm vec![0u8; 32]; hkdf.expand([], mut okm).expect(HKDF expand failed); self.pairwise_keys.insert(peer_id, okm); } } // 2. Masking掩码生成与施加 impl SecureAggClient { /// 生成 pairwise mask —— 伪随机数生成器 /// mask_{i,j} PRG(HKDF(s_{i,j}, mask)) fn generate_mask(self, peer_id: DeviceId) - Vecf32 { let key self.pairwise_keys.get(peer_id).expect(no pairwise key); // 使用 HKDF 再次派生 mask 专用密钥避免与加密密钥冲突 let hkdf Hkdf::Sha256::new(Some(bmask_derivation), key); let mut prg_seed vec![0u8; 32]; hkdf.expand([], mut prg_seed).expect(HKDF expand failed); // PRG —— 从 seed 确定性生成伪随机梯度 // 使用 seed 模型序号作为索引生成每个参数 let mut rng { let mut hasher Sha256::new(); hasher.update(prg_seed); let hash hasher.finalize(); // 从哈希创建确定性 RNG let seed: [u8; 32] hash.into(); rand::rngs::StdRng::from_seed(seed) }; // 生成与梯度同维度的随机向量 (0..1000) // 实际应为 gradient.len() .map(|_| { // 将 u32 → f32 映射到 [-1, 1] let val rng.next_u32() as f32 / u32::MAX as f32; val * 2.0 - 1.0 }) .collect() } /// 为原始梯度施加 mask /// masked_gradient_i ΔW_i Σ_{j: ji} mask_{i,j} - Σ_{j: ji} mask_{i,j} /// 当所有设备上传 masked_gradient 后pairwise mask 在聚合中抵消 pub fn mask_gradient(self, raw_gradient: Gradient, peers: [DeviceId]) - Gradient { let mut masked raw_gradient.clone(); for peer_id in peers { if peer_id self.device_id { continue; } let mask self.generate_mask(peer_id); let sign if self.device_id peer_id { 1.0f32 } else { -1.0f32 }; for (i, g) in masked.iter_mut().enumerate() { // sign 决定方向: 小 ID 加 mask大 ID 减 mask // 此约定保证所有设备聚合时 mask 成对抵消 *g sign * mask[i]; } } masked } } // 3. 差分隐私噪声注入 impl SecureAggClient { /// 向梯度注入高斯噪声 —— (ε, δ)-DP 实现 /// /// 算法: Gaussian Mechanism /// σ sqrt(2 * ln(1.25/δ)) * C / ε /// noise ~ N(0, σ²) pub fn add_dp_noise(self, gradient: mut Gradient) { // 计算高斯噪声标准差 // σ Δf * sqrt(2 * ln(1.25/δ)) / ε // 敏感度 Δf 2 * C (梯度被裁剪到 [-C, C]) let sensitivity 2.0 * self.dp_config.clip_norm as f64; let sigma sensitivity * (2.0 * (1.25 / self.dp_config.delta).ln()).sqrt() / self.dp_config.epsilon; // 生成高斯噪声 (Box-Muller 变换) let mut rng OsRng; for g in gradient.iter_mut() { let u1: f64 (rng.next_u32() as f64) / (u32::MAX as f64); let u2: f64 (rng.next_u32() as f64) / (u32::MAX as f64); // Box-Muller: z sqrt(-2 * ln(u1)) * cos(2π * u2) let z (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos(); let noise z * sigma; *g noise as f32; } } /// 梯度裁剪 —— 将单样本梯度的 L2 范数限制在 C 以内 /// 防止单个异常样本贡献过大的梯度更新 pub fn clip_gradient(self, gradient: mut Gradient) { let l2_norm: f32 gradient.iter().map(|g| g * g).sum::f32().sqrt(); if l2_norm self.dp_config.clip_norm { let scale self.dp_config.clip_norm / l2_norm; for g in gradient.iter_mut() { *g * scale; } } } /// 完整的本地处理流程: /// 1. 裁剪 → 2. 加 DP 噪声 → 3. 加 MPC mask pub async fn process_gradient( mut self, raw_gradient: mut Gradient, peers: [DeviceId], ) - Gradient { // 1. 裁剪 self.clip_gradient(raw_gradient); // 2. 添加差分隐私噪声 // 注意DP 噪声在 masking 之前添加 // 因为 mask 是用来保护传输安全的噪声是用来保护隐私的 self.add_dp_noise(raw_gradient); // 3. MPC masking self.mask_gradient(raw_gradient, peers) } } // 4. 服务端聚合 impl SecureAggServer { pub fn new() - Self { Self { public_keys: HashMap::new(), aggregated: Mutex::new(None), alive_devices: Mutex::new(Vec::new()), } } /// 更新存活设备列表 pub async fn set_alive_devices(self, device_ids: VecDeviceId) { *self.alive_devices.lock().await device_ids; } /// 聚合 masked gradients /// 当所有 mask 正确生成时pairwise mask 相互抵消 /// 服务器最终得到: Σ(ΔW_i noise_i) ΣΔW_i Σnoise_i pub async fn aggregate( self, gradients: [(DeviceId, Gradient)], ) - Gradient { if gradients.is_empty() { return vec![]; } let size gradients[0].1.len(); let mut sum vec![0.0f32; size]; let n gradients.len() as f32; for (_, gradient) in gradients { for (i, g) in gradient.iter().enumerate() { sum[i] g; } } // FedAvg: 除以设备数量取平均 for s in sum.iter_mut() { *s / n; } let mut agg self.aggregated.lock().await; *agg Some(sum.clone()); sum } } // 5. 完整协议流程示意 #[tokio::test] async fn test_secure_aggregation_flow() { let dp_config DpConfig { epsilon: 8.0, // 适中的隐私预算 delta: 1e-5, clip_norm: 1.0, }; // 3 个边缘设备 let mut client_a SecureAggClient::new(1, dp_config.clone()); let mut client_b SecureAggClient::new(2, dp_config.clone()); let mut client_c SecureAggClient::new(3, dp_config.clone()); // 密钥协商交换公钥并建立 pairwise keys let pk_a client_a.public_key(); let pk_b client_b.public_key(); let pk_c client_c.public_key(); client_a.establish_pairwise_key(2, pk_b); client_a.establish_pairwise_key(3, pk_c); client_b.establish_pairwise_key(1, pk_a); client_b.establish_pairwise_key(3, pk_c); client_c.establish_pairwise_key(1, pk_a); client_c.establish_pairwise_key(2, pk_b); // 各自训练并产生梯度 let mut grad_a vec![0.5f32; 1000]; let mut grad_b vec![0.3f32; 1000]; let mut grad_c vec![0.2f32; 1000]; // 处理梯度裁剪 DP噪声 Masking let peers vec![1, 2, 3]; let masked_a client_a.process_gradient(mut grad_a, peers).await; let masked_b client_b.process_gradient(mut grad_b, peers).await; let masked_c client_c.process_gradient(mut grad_c, peers).await; // 服务器聚合 let server SecureAggServer::new(); let aggregated server.aggregate([ (1, masked_a), (2, masked_b), (3, masked_c), ]).await; // 验证: 聚合结果应接近平均值 0.333... assert!((aggregated[0] - 0.333).abs() 1.0, DP noise adds variance but mean should be close); }关键设计决策HKDF 多层密钥派生ECDH Raw Secret → HKDF(peer sort) → Pairwise Key → HKDF(mask) → PRG Seed。每一层派生使用不同的info参数确保密钥域隔离——攻击者即使破解了 mask PRG 种子也无法推导出 pairwise key。高斯机制 σ 计算公式σ sqrt(2*ln(1.25/δ)) * Δf / ε来自 Dwork Roth, 2014。ε 取 8 时噪声适中ε 取 1 时保护极强但准确度显著下降。噪声在 masking 之前添加如果噪声在 masking 之后添加mask 的抵消逻辑会受噪声干扰。DP 噪声必须由每个设备独立施加不能在服务器端统一添加。梯度裁剪的 L2 范数裁剪上限 C 是超参数。C 太大→DP 噪声也大σ ∝ C/εC 太小→丢失有效梯度信息。通常从数据分布中取 90 百分位作为初始值。四、联邦学习安全增强的适用边界与权衡适用场景医疗、金融等隐私法规严格GDPR/HIPAA的行业。移动设备上的联邦学习Gboard 输入法预测设备数量 1000。跨组织数据协作各方既希望联合建模又不信任对方。不适用场景所有数据在同一数据中心内——直接使用中心化训练更高效。设备数量 10 的场景——安全聚合的密钥协商开销占据了训练时间的主要部分。模型极小 100 参数DP 噪声相对梯度的比例过大模型无法收敛。主要权衡安全聚合的通信开销每个设备需要与所有其他设备进行 DH 密钥交换通信复杂度 O(N²)。对于 N1000这意味着一轮训练中每个设备需要发送 999 条密钥交换消息。DP 的精度损失ε8 时准确度损失约 2-3%与任务相关ε1 时损失可达 10-15%。需要在隐私预算和模型质量之间找到平衡。设备掉线的鲁棒性安全聚合的 unmasking 阶段需要存活设备上传掉线设备的秘密份额。在最坏情况下N-1 个设备掉线唯一存活设备承担全部通信。五、总结安全聚合通过 MPC 协议的 pairwise masking保证服务器只能获得梯度总和无法窥视单设备梯度。差分隐私通过向梯度注入高斯噪声保证聚合结果不泄露单个训练样本的信息。安全聚合保护传输过程差分隐私保护聚合结果——二者正交互补不是替代关系。HKDF 多层密钥派生实现密钥域隔离是 MPC 协议安全性的基础保障。隐私预算 ε 直接决定了 DP 噪声强度与模型精度之间的取舍——是联邦学习系统的核心超参数。