联邦学习实战:FedAvg+SMOTE信用卡欺诈检测源码解析
简介这份资源面向计算机、人工智能、通信工程等专业的在校学生与算法初学者提供一套基于FedAvg联邦学习算法与SMOTE过采样优化的联邦信用卡欺诈交易检测完整项目源码。项目通过构建Server与Clients对象模拟真实场景下服务器与节点间的双向参数传递在保护各银行数据隐私、避免数据集跨机构共享的前提下缓解欺诈样本极度不平衡问题可用于毕业设计、课程设计、算法进阶学习或项目初期立项演示。压缩包共8个文件约43.14MB包含5个Python源码文件、1个Markdown说明文档、1张流程示意图与1份信用卡交易数据集分别对应模型定义、服务端与客户端逻辑、数据处理及运行说明等模块。目前已有173人学习关注。代码均经测试运行成功答辩评审平均分达96分读者可据此理解联邦学习参数聚合流程、SMOTE过采样在欺诈检测中的落地方式并在此基础上修改扩展功能。1. 联邦学习遇上信用卡欺诈一份能跑通的 FedAvg SMOTE 实战源码信用卡欺诈检测是机器学习里典型的「极端不平衡 数据不能出库」双难题。银行之间因为隐私和合规交易数据没法汇总到一处训练而单家银行的欺诈样本又少得可怜模型很容易学成「全部预测为正常」的废物。这份Federated-Learning-with-Pytorch-master源码用 FedAvg 联邦学习算法把 Server 和多个 Client 串起来模拟真实场景下服务器与节点之间的双向参数传递再叠加 SMOTE 过采样在本地把欺诈样本补足让每个客户端都能在本地学到有意义的欺诈特征。它适合正在做毕设、课设或者想入门联邦学习代码的 Python 学习者——不是纯理论科普是能直接python main.py跑起来看结果的那种。2. 拆开源码包Server/Client 架构与 SMOTE 到底怎么接进去2.1 文件清单与各自职责拿到压缩包解压后根目录下是这些文件文件作用main.py总入口负责初始化 Server、分发模型、启动联邦训练循环server.py定义 Server 类聚合各 Client 上传的模型参数FedAvg 的核心client.py定义 Client 类本地训练 SMOTE 过采样 上传参数model.py定义 PyTorch 网络结构一个简单的全连接二分类器load_data.py读取creditcard.csv做特征标准化和训练/测试划分creditcard.csv信用卡交易数据集含Class标签列0 正常 / 1 欺诈process.png训练过程可视化图方便对照结果README.md运行说明和依赖列表这个结构很干净没有多余的封装适合逐文件读。main.py是唯一需要手动执行的脚本其余都是被它 import 的模块。2.2 FedAvg 的参数传递逻辑FedAvg 的核心思想一句话Server 把全局模型下发给各 ClientClient 在本地数据上训练若干轮后把参数传回Server 按样本量加权平均更新全局模型。这份代码里Server 和 Client 之间的「双向参数传递」是通过 PyTorch 的state_dict()来做的——不是传梯度是传整个模型权重。# server.py 核心聚合逻辑示意 def aggregate(self, client_models, client_sizes): global_dict self.global_model.state_dict() total_size sum(client_sizes) for key in global_dict.keys(): # 按各客户端样本量加权平均样本多的客户端话语权更大 global_dict[key] sum( client_models[i][key] * client_sizes[i] / total_size for i in range(len(client_models)) ) self.global_model.load_state_dict(global_dict) return self.global_model这里的关键参数是client_sizes也就是每个客户端参与训练的样本数。如果某家银行数据量大它的模型更新在全局聚合时权重就高。常见做法是直接用本地训练样本总数但如果你想让各客户端更均衡也可以改成等权平均——把client_sizes[i] / total_size换成1 / len(client_models)即可。两种方式各有适用场景数据量差异大时用加权差异小时用等权更稳。2.3 SMOTE 在 Client 端的接入位置SMOTE 不能放在 Server 端做因为 Server 根本拿不到原始数据——这正是联邦学习的意义。所以过采样必须在每个 Client 的本地训练前完成。代码里client.py的train方法大致是这样组织的# client.py 本地训练 SMOTE示意 from imblearn.over_sampling import SMOTE def train(self): X, y self.local_data # 本地数据欺诈样本极少 # 只在训练集上做 SMOTE测试集保持原始分布 smote SMOTE(random_state42) X_res, y_res smote.fit_resample(X, y) # 转成 Tensor 后送入本地模型训练 X_tensor torch.tensor(X_res, dtypetorch.float32) y_tensor torch.tensor(y_res, dtypetorch.float32).unsqueeze(1) # ... 本地 epoch 循环反向传播更新 self.model return self.model.state_dict(), len(X_res)注意fit_resample只对训练数据做测试集绝对不能碰 SMOTE否则评估指标会虚高——这是血泪经验很多人第一次跑就栽在这里。random_state42是为了结果可复现你可以改成任意整数但同一组实验里要保持一致。返回的len(X_res)是过采样后的样本数Server 聚合时会用到这个值做加权。2.4 数据加载与特征处理load_data.py负责把creditcard.csv读进来。这个数据集原始特征是Time、V1~V28、Amount、Class。常见做法是丢掉Time对Amount做标准化V1~V28本身已经是 PCA 降维后的结果一般不再处理。# load_data.py 关键步骤示意 import pandas as pd from sklearn.preprocessing import StandardScaler def load_creditcard(path): df pd.read_csv(path) df df.drop(columns[Time]) # Time 对欺诈识别贡献低去掉 scaler StandardScaler() df[Amount] scaler.fit_transform(df[[Amount]]) # Amount 量纲差异大标准化 X df.drop(columns[Class]).values y df[Class].values return X, y标准化器只在训练集上fit然后transform测试集——如果你把整个数据集一起 fit就造成了数据泄漏。这份代码在划分客户端数据时我一般会建议按行切分模拟不同银行比如前 60% 给 Client 0后 40% 给 Client 1而不是随机打乱后均分因为真实场景下各银行的数据分布本来就不一样。3. 跑起来环境配置、启动命令与参数调优3.1 环境依赖与安装这份代码依赖 PyTorch、imbalanced-learn、pandas、scikit-learn、numpy。Python 版本建议 3.8~3.10太新的版本有时 imbalanced-learn 编译会出问题。用 conda 或 venv 建一个干净环境# 创建虚拟环境以 conda 为例 conda create -n fl_fraud python3.9 conda activate fl_fraud # 安装核心依赖 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install imbalanced-learn pandas scikit-learn numpy matplotlib如果你有 GPU把第一行换成对应 CUDA 版本的安装命令即可。CPU 版本跑这个规模的数据集完全够用creditcard.csv只有 28 万条左右全连接网络参数量很小。3.2 启动训练与关键参数入口是main.py直接运行python main.py但跑之前建议先打开main.py看几个参数# main.py 中常见的可调参数示意 NUM_CLIENTS 2 # 模拟几家银行 NUM_ROUNDS 10 # 联邦通信轮数 LOCAL_EPOCHS 5 # 每个客户端本地训练轮数 BATCH_SIZE 64 LR 0.001 # 学习率NUM_ROUNDS和LOCAL_EPOCHS是最影响结果的组合。轮数太少全局模型没收敛太多则通信开销大且可能过拟合。我一般先设NUM_ROUNDS10、LOCAL_EPOCHS5跑一遍看 loss 曲线如果还在下降就加到 20 轮。学习率0.001是 Adam 的常用起点如果 loss 震荡就降到0.0005。3.3 评估指标怎么看信用卡欺诈检测不能只看准确率——因为正常样本占 99.8%全预测正常也有 99.8% 准确率但一个欺诈都抓不到。要重点看召回率Recall和 F1 分数以及 AUC-ROC。# 评估部分示意 from sklearn.metrics import classification_report, roc_auc_score y_pred (model(X_test) 0.5).float() print(classification_report(y_test, y_pred)) print(AUC:, roc_auc_score(y_test, y_pred.detach().numpy()))跑完后对照process.png里的曲线如果你的召回率明显低于图里展示的水平大概率是 SMOTE 没生效或者测试集被过采样污染了。先检查client.py里fit_resample的调用位置再确认测试集是否独立。3.4 调整客户端数量模拟不同场景想模拟更多银行参与改NUM_CLIENTS就行但数据切分逻辑也要跟着改。常见做法是用numpy.array_split把数据均分# 按客户端数量切分数据示意 import numpy as np indices np.array_split(np.arange(len(X)), NUM_CLIENTS) client_data [(X[idx], y[idx]) for idx in indices]注意欺诈样本本身很少切分后某些客户端可能一个欺诈样本都没有SMOTE 会直接报错。解决办法是先做一次分层切分保证每个客户端至少有几个欺诈样本或者把SMOTE的k_neighbors参数调小默认 5样本太少时改成 1 或 2。4. 避坑与排查跑不通、指标异常、显存爆了怎么办4.1 SMOTE 报错 Expected n_neighbors n_samples现象运行到 Client 本地训练时抛出ValueError: Expected n_neighbors n_samples, but n_samples 3, n_neighbors 6。原因某个客户端的欺诈样本数少于 SMOTE 默认的k_neighbors5无法构造合成样本。解决在SMOTE()里显式指定k_neighbors1或者先检查各客户端欺诈样本数样本太少的客户端直接跳过过采样、用原始数据训练。我一般会在切分数据后打印一句print(fClient {i} fraud samples: {sum(y)})心里有数再跑。4.2 测试集指标高得离谱现象召回率 0.99F1 接近 1.0但换一组数据就崩。原因SMOTE 被错误地应用到了测试集或者标准化器在整个数据集上 fit 造成了数据泄漏。解决确认fit_resample只在训练集调用标准化器先fit训练集再transform测试集。这两步是铁律没有例外。4.3 全局模型不收敛loss 来回震荡现象每轮聚合后 loss 忽高忽低准确率上不去。原因各客户端数据分布差异太大Non-IIDFedAvg 简单加权平均无法调和或者学习率过高。解决先把学习率降到0.0001试一轮如果还震荡考虑增加LOCAL_EPOCHS让各客户端本地充分收敛后再上传或者改用按样本量加权的聚合方式。极端 Non-IID 场景下 FedAvg 本身就有局限这是算法边界不是代码 bug。4.4 CUDA out of memory现象有 GPU 但一跑就爆显存。原因creditcard.csv虽然不大但如果 batch size 设得太大或者同时开了多个客户端并行训练显存会不够。解决把BATCH_SIZE从 64 降到 32 或 16确保客户端是串行训练而不是并行。CPU 跑这个数据集完全可行不必强求 GPU。4.5 依赖版本冲突导致 import 失败现象import imblearn报错或者 PyTorch 和 numpy 版本不兼容。原因imbalanced-learn 对 scikit-learn 版本有要求numpy 2.x 和旧版 PyTorch 也可能冲突。解决用pip install imbalanced-learn0.11.0 scikit-learn1.3.0 numpy1.24.0锁定版本。如果还不行建一个全新的 conda 环境从头装别在旧环境里折腾。5. 进阶技巧把 FedAvg 换成 FedProx、验证 SMOTE 是否真的有用5.1 用消融实验验证 SMOTE 的贡献很多人跑完不知道 SMOTE 到底有没有用。最直接的办法是做一组对照把client.py里的fit_resample注释掉其他参数不变再跑一次对比召回率。# 消融实验关闭 SMOTE # X_res, y_res smote.fit_resample(X, y) # 注释掉这行 X_res, y_res X, y # 直接用原始数据如果关闭 SMOTE 后召回率从 0.85 掉到 0.3 以下说明过采样确实在起作用如果差别不大可能是你的模型容量不够或者学习率没调好SMOTE 补出来的样本没被有效利用。我一般会把两组结果的classification_report并排贴出来看比只看一个数字靠谱得多。5.2 从 FedAvg 迁移到 FedProxFedAvg 在 Non-IID 数据上容易发散FedProx 通过在本地损失里加一个近端项来约束客户端模型不要偏离全局模型太远。改动很小在client.py的损失函数里加一项# FedProx 近端项示意 mu 0.01 # 近端项系数控制约束强度 global_params [p.clone().detach() for p in global_model.parameters()] proximal_term sum( ((local_param - global_param) ** 2).sum() for local_param, global_param in zip(model.parameters(), global_params) ) loss criterion(output, target) (mu / 2) * proximal_termmu是关键参数设 0 就退化成 FedAvg设太大本地模型学不动。常见做法是从0.01开始试观察 loss 曲线是否比 FedAvg 更平滑。这个改动不需要动 Server 端聚合逻辑完全复用。5.3 一个验证聚合是否正确的笨办法联邦学习代码最容易出错的地方是参数聚合——传错了 key、维度对不上、或者聚合后忘了load_state_dict。我习惯在每轮聚合后加一句检查# 聚合后验证全局模型参数是否真的变了 before [p.clone() for p in global_model.parameters()] server.aggregate(client_models, client_sizes) after [p for p in global_model.parameters()] diff sum((b - a).abs().sum().item() for b, a in zip(before, after)) print(fRound {r} param diff: {diff:.6f})如果diff是 0说明聚合没生效大概率是load_state_dict没调用或者传了错误的 dict。这个检查花不了几行代码但能省掉大量「模型怎么不收敛」的排查时间。从那以后我每次跑联邦学习代码都强制先跑一轮NUM_ROUNDS1看参数 diff 和 loss 是否正常确认链路通了再放开轮数。希望帮到你。本文还有配套的精品资源点击获取