
1. 项目概述为什么 FedAvg 是联邦学习落地的“第一块砖”如果你刚接触联邦学习大概率会发现几乎所有入门教程、论文综述甚至工业界白皮书里第一个出现的算法名字就是 FedAvg——全称 Federated Averaging。它不像某些前沿变体那样挂着“自适应”“异构鲁棒”“差分隐私增强”等炫酷前缀但它稳、快、好懂、易复现更重要的是——它真正在真实设备上跑得起来。我带过三届校企联合实验室的学生从医疗影像协作建模到银行间反欺诈模型共建只要客户说“先搭个能跑通的 baseline”我们第一行代码写的永远是 FedAvg。它不是最完美的但它是唯一一个能让数据不出域、模型还能协同进化的“最小可行解”。核心关键词联邦学习、FedAvg、pytorch这三个词串起来本质上是在回答一个现实问题当数据被物理隔离在手机、医院服务器、工厂边缘设备里我们如何让它们“一起学”而不是“各自学完再拼凑”FedAvg 的答案很朴素不传原始数据只传模型更新也就是梯度或参数差然后在服务端做加权平均。这个思路看似简单背后却卡着三个硬骨头一是本地训练轮数E怎么设才不至于让手机端训太久掉电关机二是客户端采样比例C怎么平衡效率与偏差比如每次只选 10% 的活跃设备会不会漏掉关键数据分布三是模型聚合时要不要做归一化、要不要过滤异常更新——这些细节PyTorch 本身不提供现成接口全靠你手动抠。所以这篇内容不是讲“FedAvg 是什么”而是带你亲手把论文里的伪代码变成能在自己笔记本上跑通、能改参数、能看 loss 曲线、能导出模型的完整 PyTorch 工程。无论你是刚学完《PyTorch 深度学习实战》第 5 章的在校生还是正被甲方催着交 federated baseline 的算法工程师只要你装好了 PyTorch哪怕只是 CPU 版就能跟着一步步走完。它不依赖 GPU 集群不强制要求 Docker 或 Kubernetes甚至不用碰 Linux 命令行——Windows Anaconda PyCharm 就够了。接下来所有内容都基于真实调试日志、失败截图和反复重跑的收敛曲线没有“理论上可以”只有“我试过这样写才不报错”。2. 核心设计逻辑FedAvg 不是“平均主义”而是带约束的协同优化2.1 FedAvg 的数学本质从 SGD 到分布式局部更新的妥协很多人初看 FedAvg 公式会觉得它就是把多个客户端的模型参数 w_i 直接求平均w^{t1} \frac{1}{K}\sum_{k1}^K w_k^{t1}。这其实是严重误解。真正的 FedAvg 过程是服务端下发全局模型 w^t → 各客户端用本地数据执行 E 轮本地 SGD → 得到更新后模型 w_k^{tE} → 上传 w_k^{tE}或更常见的 Δw_k w_k^{tE} - w^t→ 服务端加权聚合。关键点在于它跳过了传统分布式 SGD 中每步都要同步梯度的通信开销用“本地多步训练 稀疏上传”换来了通信效率。但代价是什么是本地模型在 E 轮内会偏离全局最优方向尤其当各客户端数据非独立同分布Non-IID时——比如 A 手机全是猫图B 手机全是狗图本地训 10 轮后A 的模型头已经极度偏向猫分类B 的则死磕狗特征直接平均会导致全局模型在两类上都变弱。这就是为什么 FedAvg 在论文中明确强调E 不能太大否则灾难性遗忘catastrophic forgetting会真实发生——不是概念术语而是你跑实验时看到 test accuracy 从 85% 掉到 62% 的具体数字。我实测过在 CIFAR-10 Non-IID 划分下每个客户端仅含 2 类E1 时最终准确率 79.3%E5 时跌到 71.6%E20 直接崩到 58.1%。所以 FedAvg 的“平均”本质是在通信成本与模型漂移之间找平衡点的工程决策而非数学上的最优解。这也是为什么后续所有 FedXXX 变体几乎都在解决同一个问题怎么让本地多步训练不那么“放飞自我”有的加正则项FedProx有的动态调学习率SCAFFOLD有的做梯度矫正FedNova。但 FedAvg 本身就是那个必须先立住的“锚点”。2.2 为什么选 PyTorch 而非 TensorFlow 或 JAX当前2024 年联邦学习框架生态里TensorFlow FederatedTFF文档最全JAX 的 DP-Fed 更适合差分隐私研究但真正让工业界快速落地的还是 PyTorch。原因很实际第一PyTorch 的动态图机制让调试本地训练过程像调试普通脚本一样直观——你可以随时 print(model.fc2.weight.grad) 看梯度值而 TFF 的声明式 API 一旦报错堆栈信息常指向内部编译器新手根本无从下手第二PyTorch 生态对移动端、嵌入式支持更成熟TorchScript 和 ONNX 导出流程稳定这意味着你今天在笔记本上跑通的 FedAvg明天就能部署到 Android 手机的 TFLite 解释器里只需加一层轻量 wrapper第三也是最关键的一点PyTorch 的 nn.Module 和 optim.Optimizer 完全可控你可以精确干预每一次参数更新。比如 FedAvg 要求客户端上传的是“更新量”Δw而不是完整模型 w这就需要你在本地 optimizer.step() 后手动计算 w_new - w_old。在 PyTorch 里这三行代码就搞定old_params {name: param.data.clone() for name, param in model.named_parameters()} # ... 执行 E 轮训练 ... new_params {name: param.data for name, param in model.named_parameters()} delta {name: new_params[name] - old_params[name] for name in old_params}而在 TFF 里你得绕进tff.learning.build_federated_averaging_process的底层 builder改源码都不一定行。所以选择 PyTorch不是因为“它更流行”而是因为它把控制权交还给了开发者——当你需要在客户端插入梯度裁剪、添加噪声、或实现偏置压缩如热词里提到的“传输经过压缩的本地更新数据”PyTorch 给你留了完整的 hook 接口。2.3 FedAvg 的四大可调参数及其物理意义FedAvg 论文里只定义了 4 个核心超参但每个都直击落地痛点绝非随便设个值就能跑C客户端采样率每轮从 K 个注册客户端中随机选取 C×K 个参与训练。设 C0.1 意味着每轮只让 10% 的设备干活。很多教程直接写np.random.choice(clients, int(0.1*len(clients)))但忽略了一个现实手机/边缘设备在线状态是强波动的。我曾在一个医疗 IoT 项目里发现凌晨 2 点只有 3% 的监护仪在线而早 8 点高达 67%。如果固定 C0.1凌晨那轮可能只采到 2 台设备导致聚合方差极大。解决方案是设置最小采样数min_clients_per_round5而非固定比例。E本地训练轮数客户端用本地数据训练多少轮。它和设备算力强相关。我在测试华为 Mate 50骁龙8 Gen1跑 ResNet-18 时E5 耗时约 42 秒换成低端红米 Note 9Helio G85同样 E5 要 118 秒。如果甲方要求“单次上传耗时 60 秒”你就必须把 E 从 5 降到 2并相应增加全局轮数 T 来补偿。B本地批量大小客户端每次喂给模型的数据量。注意B 不是越大越好。低端设备内存有限B64 可能直接 OOM但 B 太小如 B4又会让梯度噪声过大影响收敛。我的经验是先用torch.cuda.memory_allocated()GPU或psutil.virtual_memory().usedCPU测设备内存余量再按公式B_max ≈ available_memory / (model_params × 4 bytes)估算上限float32 占 4 字节然后取整数倍。η客户端学习率这是最容易被忽视的“隐形杀手”。FedAvg 默认用和服务端相同的 η但实践中客户端数据量少、分布偏η 过大会导致本地训练震荡。我在一个金融风控场景中将客户端 η 从 0.01 降到 0.003test F1-score 提升了 4.7 个百分点——因为小学习率让本地模型更“谨慎”减少了 Non-IID 下的过拟合。提示这四个参数不是孤立的。E 和 B 共同决定单次本地训练耗时C 和 η 影响全局收敛速度。没有“万能配置”必须针对你的硬件清单和数据分布做网格搜索。我通常先固定 C0.1、B32、η0.01只调 E∈{1,3,5}跑 50 轮看 loss 曲线平滑度再微调其他参数。3. PyTorch 实现详解从零构建可调试、可监控、可扩展的 FedAvg3.1 环境准备与依赖确认避开 Anaconda 和 PyTorch 的经典坑别跳过这一步。我见过太多人卡在环境上不是代码问题而是版本冲突。FedAvg 对 PyTorch 版本其实不敏感但必须避开两个雷区一是 PyTorch 1.8因为torch.compile()和torch.amp在旧版不稳定影响后续加 FP16 压缩二是 PyTorch 2.2 且 CUDA 版本不匹配会导致DistributedDataParallel初始化失败。我的黄金组合是PyTorch 2.1.0 CUDA 11.8对应 NVIDIA 驱动 ≥ 520。安装命令不是简单pip install torch而是# Windows 用户推荐 pip3 install torch2.1.0cu118 torchvision0.16.0cu118 torchaudio2.1.0 --extra-index-url https://download.pytorch.org/whl/cu118 # Linux/Mac 用户若用 CPU 版替换 cu118 为 cpu pip3 install torch2.1.0cpu torchvision0.16.0cpu torchaudio2.1.0 --extra-index-url https://download.pytorch.org/whl/cpu为什么指定 2.1.0因为 2.2.0 在torch.nn.utils.clip_grad_norm_中引入了新行为会导致 FedAvg 的梯度裁剪逻辑失效——你本地训完发现 Δw 异常大查半天才发现是 PyTorch 自己改了 clip 规则。Anaconda 方面务必创建干净环境conda create -n fedavg_env python3.9 conda activate fedavg_env # 再执行上面的 pip install不要用 base 环境base 里可能有旧版 numpy 或 scipy和 PyTorch 的 BLAS 库冲突表现为你跑 FedAvg 时 CPU 占用 100% 却不训练——这是 OpenBLAS 线程死锁重装环境是最快解法。VSCode Anaconda 配合时在.vscode/settings.json里加{ python.defaultInterpreterPath: ./envs/fedavg_env/bin/python, python.testing.pytestArgs: [tests/] }确保调试器用对解释器。这些细节看着琐碎但能帮你省下至少 8 小时的无效 debug 时间。3.2 数据划分与 Non-IID 模拟不做假数据就做不出真效果FedAvg 的价值恰恰体现在 Non-IID 场景下。如果你用torch.utils.data.random_split把 MNIST 均匀分给 100 个客户端那 FedAvg 和中心化训练几乎没区别——因为每个客户端数据分布一致本地训再多轮也不会漂移。真正的挑战是模拟现实手机用户拍照偏好不同医院病种分布不均工厂传感器故障模式各异。我采用Dirichlet 分布划分法这是目前最贴近真实的 Non-IID 模拟方式。原理很简单为每个客户端 k 生成一个概率向量 p_k ∈ R^CC 是类别数满足 ∑p_k,c 1且 p_k,c ~ Dir(α)α 是控制“偏斜程度”的超参。α 越小分布越偏比如 α0.1 时一个客户端可能 90% 是数字 010% 是数字 1另一个全是数字 5α10 时接近均匀。PyTorch 实现只需 15 行def dirichlet_split_noniid(train_labels, alpha, n_clients): train_labels: torch.Tensor, shape [N] 返回: client_data_idx: List[List[int]], 每个子列表是该客户端拥有的样本索引 n_classes train_labels.max() 1 class_idxs [torch.where(train_labels i)[0] for i in range(n_classes)] client_data_idx [[] for _ in range(n_clients)] for k in range(n_classes): # 对第 k 类样本按 Dirichlet 分配给各客户端 proportions np.random.dirichlet([alpha] * n_clients) class_len len(class_idxs[k]) # 累计分配避免浮点误差 cum_sum np.cumsum(proportions) * class_len cum_sum cum_sum.astype(int) # 分割索引 start 0 for i in range(n_clients): end cum_sum[i] if i n_clients-1 else class_len client_data_idx[i].extend(class_idxs[k][start:end].tolist()) start end return client_data_idx调用时client_data_idx dirichlet_split_noniid(y_train, alpha0.5, n_clients100)。α0.5 是我的默认起点——足够体现 Non-IID又不至于让某些客户端完全缺失某类样本否则训练会崩溃。你可以在train.py开头加一行print(fClient 0 data: {Counter(y_train[client_data_idx[0]])})亲眼看到客户端 0 拥有 823 张数字 3却只有 7 张数字 8。这种真实感是理解 FedAvg 必要性的第一步。3.3 客户端本地训练模块不只是跑 E 轮 SGD客户端代码不是“加载数据→训练→上传”而是包含五个关键环节缺一不可模型初始化与状态同步客户端首次启动时必须从服务端下载初始模型 w^0。但后续轮次它收到的是 w^t需用load_state_dict()加载。注意strictFalse可能掩盖层名不匹配错误务必设strictTrue并捕获异常。数据加载器构建不能直接用DataLoader(dataset, shuffleTrue)因为 Non-IID 下 shuffle 会破坏本地数据分布特性。正确做法是shuffleFalse并在每个 epoch 开始时用torch.randperm(len(dataset))生成新索引再SubsetRandomSampler——这样既保证每个 epoch 内部打乱又保持客户端间分布差异。本地训练循环核心是 E 轮迭代。但必须加入 early stopping如果某客户端在第 3 轮 loss 就不再下降比如连续 2 轮 Δloss 1e-5就提前退出避免无效计算。代码片段for local_epoch in range(E): model.train() for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() optimizer.step() # 每轮后评估 val_acc evaluate(model, val_loader, device) if local_epoch 0 and abs(val_acc - prev_val_acc) 1e-4: break # 收敛了提前退出 prev_val_acc val_acc更新量 Δw 计算与压缩这才是 FedAvg 的灵魂。不是上传model.state_dict()而是计算差值# 训练前保存初始参数 init_state {name: param.data.clone() for name, param in model.named_parameters()} # 训练后 final_state {name: param.data for name, param in model.named_parameters()} delta_state {} for name in init_state: delta_state[name] final_state[name] - init_state[name] # 可选添加 Top-k 稀疏化对应热词“偏置压缩” if sparsity 0: for name in delta_state: k int(delta_state[name].numel() * sparsity) values, indices torch.topk(delta_state[name].abs().flatten(), k) mask torch.zeros_like(delta_state[name]).flatten() mask[indices] 1 mask mask.reshape(delta_state[name].shape) delta_state[name] delta_state[name] * masksparsity0.9 表示只传最大的 10% 更新量通信量减少 90%实测在 CIFAR-10 上精度损失 1.5%。异常处理与心跳上报客户端必须实现超时机制。如果训练超过 120 秒没响应主动断连并上报 error code。服务端据此剔除该客户端避免拖慢全局进度。3.4 服务端聚合逻辑加权平均不是简单求和服务端聚合看似简单但藏着三个易错点权重计算必须基于数据量不能w^{t1} mean([w_k^{tE} for k in selected])而应w^{t1} sum(n_k / N_total * w_k^{tE})其中 n_k 是客户端 k 的本地样本数N_total 是本轮所有选中客户端的样本总数。否则一个拥有 10000 张图的大医院和一个只有 50 张图的小诊所对全局模型的影响一样大显然不合理。参数名必须严格对齐如果客户端 A 用nn.Linear(10,5)B 用nn.Linear(10,5,biasFalse)它们的 state_dict 键名不同A 有biasB 没有直接torch.stack()会报错。解决方案是服务端定义统一模型结构客户端必须继承该结构或在上传前做键名映射。聚合后必须做梯度裁剪即使客户端做了裁剪聚合后的 Δw 仍可能爆炸。我在一个 NLP 任务中发现某客户端因文本长度异常Δw 的 L2 norm 达到 1e6直接拉垮全局模型。因此服务端聚合后必须# 聚合得到 avg_delta for name in avg_delta: torch.nn.utils.clip_grad_norm_(avg_delta[name], max_norm1.0) # 再应用到全局模型 for name in global_model.named_parameters(): if name in avg_delta: global_model.state_dict()[name].add_(avg_delta[name])完整服务端聚合函数如下已通过 100 客户端压力测试def server_aggregate(global_model, client_deltas, client_data_sizes): client_deltas: List[Dict[str, torch.Tensor]], 每个字典是客户端上传的 delta_state client_data_sizes: List[int], 对应客户端的样本数 total_size sum(client_data_sizes) # 初始化 avg_delta avg_delta {name: torch.zeros_like(param.data) for name, param in global_model.named_parameters()} # 加权累加 for delta, size in zip(client_deltas, client_data_sizes): weight size / total_size for name in avg_delta: if name in delta: avg_delta[name] weight * delta[name] # 全局裁剪 for name in avg_delta: torch.nn.utils.clip_grad_norm_(avg_delta[name], max_norm1.0) # 应用更新 with torch.no_grad(): for name, param in global_model.named_parameters(): if name in avg_delta: param.add_(avg_delta[name]) return global_model3.5 完整训练流程与监控让 FedAvg “看得见、摸得着”一个无法监控的 FedAvg 是危险的。我坚持在每轮训练后输出四维指标全局指标global_test_acc,global_test_loss在中心化测试集上评估客户端指标均值mean_client_train_acc,std_client_train_acc100 个客户端训练准确率的均值和标准差。如果 std 15%说明 Non-IID 严重需调小 E 或加正则。通信统计total_upload_bytes累计上传字节数avg_client_time_sec客户端平均耗时。这是验证“偏置压缩”是否有效的直接证据。模型健康度grad_norm_global全局模型梯度 L2 normweight_variance各层参数方差。如果grad_norm_global连续 5 轮 1e-5说明模型已饱和可提前终止。监控代码集成在主循环里for round_idx in range(T): # 1. 采样客户端 selected_clients sample_clients(all_clients, C) # 2. 并行训练这里用 threading 简化生产用 asyncio client_results [] for client in selected_clients: result client.local_train(global_model, E, B, eta) client_results.append(result) # 3. 聚合 client_deltas [r[delta] for r in client_results] client_sizes [r[data_size] for r in client_results] global_model server_aggregate(global_model, client_deltas, client_sizes) # 4. 监控输出 global_acc, global_loss evaluate(global_model, test_loader, device) client_accs [r[train_acc] for r in client_results] print(f[Round {round_idx}] fGlobal Acc: {global_acc:.4f} | fClient Acc Mean±Std: {np.mean(client_accs):.4f}±{np.std(client_accs):.4f} | fUpload MB: {sum(r[upload_mb] for r in client_results):.2f} | fTime Avg: {np.mean([r[time_sec] for r in client_results]):.2f}s) # 5. 保存 checkpoint if round_idx % 10 0: torch.save(global_model.state_dict(), fcheckpoints/global_round_{round_idx}.pth)运行后你会看到类似这样的日志[Round 0] Global Acc: 0.1243 | Client Acc Mean±Std: 0.1187±0.0214 | Upload MB: 12.45 | Time Avg: 38.21s [Round 10] Global Acc: 0.4521 | Client Acc Mean±Std: 0.4432±0.0876 | Upload MB: 124.5 | Time Avg: 41.03s [Round 50] Global Acc: 0.7892 | Client Acc Mean±Std: 0.7765±0.1243 | Upload MB: 623.1 | Time Avg: 42.17s注意Client Acc Mean±Std这一项从 ±0.0214 到 ±0.1243说明随着训练进行客户端间性能差距在拉大——这正是 Non-IID 的典型表现也是 FedAvg 需要持续优化的原因。没有这个监控你就像蒙着眼开车。4. 实战问题排查那些论文里不会写的“血泪教训”4.1 问题速查表高频报错与根因定位报错信息根本原因解决方案我的实测耗时RuntimeError: Expected all tensors to be on the same device客户端模型在 CPU服务端在 GPU或反之统一设备管理定义device torch.device(cuda if torch.cuda.is_available() else cpu)所有 tensor 创建时显式指定device2 分钟KeyError: bias客户端模型结构与服务端不一致如某层少了 bias服务端定义BaseModel(nn.Module)所有客户端继承它或在server_aggregate前加键名校验assert set(avg_delta.keys()) set(global_model.state_dict().keys())5 分钟CUDA out of memory本地批量 B 过大或模型太深用torch.cuda.memory_summary()查内存瓶颈降低 B或用torch.compile(model)优化显存15 分钟需重跑NaN loss during training学习率 η 过大或数据未归一化检查输入数据torch.isnan(data).any()η 从 0.001 开始试加torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)8 分钟Global accuracy stuck at 0.1Non-IID 过于严重α0.1或 E 过大用dirichlet_split_noniid时增大 α 至 0.5减小 E 至 1或在客户端加torch.nn.CrossEntropyLoss(label_smoothing0.1)20 分钟需重新划分数据这张表来自我过去三年 17 个联邦学习项目的 debug 日志。特别提醒NaN loss问题90% 源于数据预处理。比如你用transforms.Normalize((0.5,),(0.5,))处理灰度图但客户端传来的图是 RGB 三通道均值维度不匹配就会在output model(data)后立刻产生 NaN。解决方案不是改模型而是加断言def validate_input(data): assert data.dim() 4, fInput must be 4D, got {data.dim()} assert data.shape[1] in [1,3], fChannel must be 1 or 3, got {data.shape[1]} assert not torch.isnan(data).any(), Input contains NaN return data放在DataLoader的collate_fn里一劳永逸。4.2 “灾难性遗忘”的实证分析与缓解策略热词里提到“灾难性遗忘 联邦学习”这不是玄学术语而是你能用 TensorBoard 看到的曲线。我在一个跨医院皮肤癌诊断项目中记录了客户端 A三甲医院数据丰富和客户端 B社区诊所仅 200 张图在 FedAvg 下的表现第 1-10 轮A 的本地 test acc 从 65% → 82%B 从 42% → 58%第 11-20 轮A 稳定在 81-83%B 却从 58% → 49% → 41%第 21 轮B 的模型在“黑色素瘤”类上准确率暴跌至 12%而“脂溢性角化”类升至 89%这就是灾难性遗忘B 的模型为了适配全局更新过度优化了常见类牺牲了其专精的罕见病种。根因是 FedAvg 的全局平均强制 B 的模型向 A 靠拢而 B 的数据量太少无法抵抗这种拉扯。缓解方法有三客户端正则化FedProx在客户端 loss 上加 proximal termL_prox L_local μ/2 ||w - w^t||^2。μ0.1 是我的起点它让本地训练“别离全局模型太远”。代码只需改一行# 原 loss loss criterion(output, target) # 加 proximal term prox_term 0 for name, param in model.named_parameters(): prox_term torch.norm(param - global_state[name]) ** 2 loss 0.1 * prox_term个性化头部Per-Head全局共享 backbone每个客户端训练自己的 classifier head。这样 B 可以保留对罕见病的判别能力。实现上服务端只聚合 backbone 参数head 参数不上传。动态权重调整服务端根据客户端历史贡献如上传 Δw 的信噪比动态调整聚合权重而非固定按数据量加权。我用weight_k exp(-mse(delta_k, avg_delta) / temperature)temperature0.5 效果最好。实测结果加 FedProx 后B 的黑色素瘤类准确率从 12% 回升至 63%全局 acc 仅降 0.4 个百分点。这证明灾难性遗忘不是 FedAvg 的缺陷而是使用方式的问题。4.3 通信开销优化实录“偏置压缩”到底能省多少热词强调“传输经过压缩的本地更新数据来减少通信开销”我用真实数据告诉你效果。在 ResNet-1811M 参数上对比三种压缩策略策略上传字节数单客户端全局 acc50轮相对节省原始 float3244 MB78.92%0%Top-k 稀疏k10%4.4 MB77.65%90%Quantizationint811 MB78.31%75%Top-k int81.1 MB76.89%97.5%关键发现Top-k 稀疏对通信量削减最狠但精度损失也最大int8 量化更温和且硬件友好NPU 直接支持。生产环境我推荐组合策略对卷积层用 int8对全连接层用 Top-k因为 FC 层参数更稀疏。实现 int8 量化只需两行# 量化 delta_int8 torch.quantize_per_tensor(delta_float, scale0.1, zero_point0, dtypetorch.qint8) # 反量化服务端聚合前 delta_float delta_int8.dequantize()scale0.1 是经验值可通过delta_float.abs().max() / 127动态计算。这样一个 11MB 的模型更新压到 1.1MB意味着在 1Mbps 上传带宽下单次上传从 88 秒降到 8.8 秒——这才是“减少通信开销”的真实意义。4.4 跨平台部署避坑指南从笔记本到安卓手机最后分享一个硬核经验FedAvg 代码写完只是开始。真正落地要过三关Windows/Linux/macOS 兼容路径分隔符用os.path.join()不要硬写/文件读写加encodingutf-8multiprocessing在 Windows 需if __name__ __main__:保护。CPU/GPU 自适应不要写model.cuda()而用model.to(device)device 由torch.device(cuda if torch.cuda.is_available() else cpu)动态决定。安卓端部署PyTorch Mobile 要求模型用 TorchScript。在服务端保存时# 训练完 traced_model torch.jit.trace(global_model, example_input) traced_model.save(fedavg_model.pt)example_input 是(1,3,224,224)这样的 dummy input。客户端Android用Module.load(fedavg_model.pt)加载比 Python 解释器快 3 倍内存占用低 60%。我曾把 FedAvg 模型部署到一台 2018 款红米 Note 74GB RAM骁龙660执行一次本地训练E3B16耗时 23.4 秒内存峰值 1.2GB完全可用。这证明联邦学习不是实验室玩具而是能跑在真实设备上的技术。5. 进阶思考FedAvg 之后路在何方FedAvg 是起点不是终点。我在实际项目中从不把它当最终方案而是作为 baseline 和调试基线。它的局限性清晰可见对 Non-IID 敏感、无法处理系统异构设备算力差异、