联邦学习+知识蒸馏实战:FL-IDS-KD入侵检测源码解析与避坑指南
2026/9/24 18:18:20 网站建设 项目流程

简介:本资源面向计算机、人工智能、通信工程等专业的在校学生与研究人员,提供一套将联邦学习与知识蒸馏结合用于网络入侵检测的完整Python实现,并在NSL-KDD数据集上完成验证,适合作为毕设、课程设计或项目立项的参考方案。压缩包共63个文件,约26.18MB,以12个py源码文件为核心,辅以pyc编译文件、txt说明、log日志、weight模型权重及png对比图等,涵盖服务端、客户端、模型定义、参数配置与GUI交互等模块。目前已有231人学习下载。读者可获取可运行的联邦训练流程、知识蒸馏建模思路、NSL-KDD数据处理脚本以及结果对比图,便于理解多客户端协同训练与模型压缩的落地方式,并在此基础上修改扩展功能。

1. 联邦学习遇上知识蒸馏:这套入侵检测源码到底能跑出什么

很多人第一次看到「联邦学习 + 知识蒸馏 + 入侵检测」这三个词堆在一起,第一反应是论文缝合怪。但我把FL-IDS-KD-master.zip解压跑通之后发现,它其实解决了一个很现实的问题:你手上有 NSL-KDD 这种标注好的流量数据,但真实场景里数据分散在不同节点、不能集中上传,而单个节点数据量又不够训出高精度模型。这套源码用联邦学习让多个客户端各自训练本地模型,再用知识蒸馏把大模型的能力压缩到小模型上,最终在 NSL-KDD 上验证了检测效果。它适合做毕设、课设、安全方向入门的人,也适合想搞懂联邦学习代码到底怎么落地的人。下面我按「是什么 → 怎么跑 → 坑在哪 → 怎么改」的顺序拆一遍。

2. 先搞懂架构再动手:FL-IDS-KD 的模块划分与数据流

2.1 服务端与客户端到底各干什么

这套代码的目录结构分得很清楚,server文件夹和client文件夹各自独立,中间靠 socket 通信。服务端main_server.py负责聚合各客户端上传的模型参数,客户端main_client_1.pymain_client_2.py负责本地训练和上传。model.py定义神经网络结构,utils.py放数据加载和预处理函数,argu.py集中管理超参数,connFun.py封装 socket 连接逻辑。initDate.py做初始化,GUI.py提供图形界面。

联邦学习的核心逻辑在服务端聚合那一步。每个客户端训练完一轮后,把local_model.weight上传到服务端,服务端按样本量加权平均得到全局模型,再下发回客户端。知识蒸馏体现在客户端本地:用一个较大的教师模型指导较小的学生模型训练,学生模型才是最终上传的那个。这样做的好处是通信量小,因为学生模型参数量少,上传下载都快。

NSL-KDD 数据集放在data目录下,temp.csvdata3.log是训练过程中生成的中间文件。resultCompare1.pngresultCompare2.png是跑完之后生成的对比图,能直观看到联邦学习和单独训练的效果差异。

2.2 环境准备与依赖安装

这套代码是 Python 3.6 环境下写的,__pycache__里的.pyc文件后缀是cpython-36,说明作者当时用的就是 3.6。我建议你用 Python 3.6 到 3.8 之间的版本,太新的版本某些库可能不兼容。依赖库主要是 PyTorch、NumPy、Pandas、Matplotlib,GUI 用的是 tkinter,Python 自带。

先建虚拟环境,避免污染全局:

python -m venv fl_ids_env source fl_ids_env/bin/activate # Linux/Mac # fl_ids_env\Scripts\activate # Windows

然后安装核心依赖。PyTorch 版本别装太新,1.4 到 1.8 之间比较稳:

pip install torch==1.8.0 torchvision==0.9.0 pip install numpy pandas matplotlib scikit-learn

如果你用的是 Windows,tkinter 一般自带,不用额外装。Linux 下如果报No module named tkinter,用sudo apt-get install python3-tk补上。

2.3 数据准备与预处理流程

NSL-KDD 数据集需要放在data目录下。原始数据是KDDTrain+.txtKDDTest+.txt,代码里的utils.py会做几件事:把符号特征做 one-hot 编码,数值特征做归一化,标签列转成二分类或多分类。initDate.py负责把原始数据切成训练集和测试集,并按照联邦学习的设定分给不同客户端。

我一般会先单独跑一下数据加载,确认没问题再启动训练:

from utils import load_data, preprocess # 加载原始数据 train_data, test_data = load_data('./data/KDDTrain+.txt', './data/KDDTest+.txt') # 预处理:归一化 + one-hot X_train, y_train, X_test, y_test = preprocess(train_data, test_data) print('训练集形状:', X_train.shape) print('测试集形状:', X_test.shape) print('类别分布:', y_train.value_counts().to_dict())

这段代码的逻辑是先把原始 txt 读进来,preprocess函数内部会对protocol_typeserviceflag这三个符号特征做独热编码,对src_bytesdst_bytes等数值特征做 Min-Max 归一化。标签列label会被映射成 0 和 1,0 表示正常流量,1 表示攻击流量。跑完打印形状,正常应该是几万条训练样本、十几万条测试样本。如果形状不对,检查数据文件路径和分隔符,NSL-KDD 原始文件是逗号分隔的。

2.4 启动服务端与多客户端训练

运行顺序很关键:先起服务端,再起客户端。服务端main_server.py会监听指定端口,等待客户端连接。客户端默认开两个窗口,分别跑main_client_1.pymain_client_2.py

服务端启动命令:

python main_server.py

客户端启动命令,开两个终端分别执行:

python main_client_1.py python main_client_2.py

argu.py里可以改几个关键参数:num_rounds控制联邦学习轮数,默认可能是 10 或 20;local_epochs控制客户端本地训练轮数;lr是学习率;batch_size是批大小。我建议第一次跑先把num_rounds设小一点,比如 5,确认流程通了再加大。

服务端和客户端之间的通信走 socket,connFun.py里封装了send_msgrecv_msg函数。如果你在本地跑,IP 用127.0.0.1就行。如果要在局域网内多机跑,把服务端 IP 改成实际地址,客户端连接时填服务端的 IP。

2.5 GUI 界面的使用方式

GUI.py提供了一个简单的图形界面,适合不习惯命令行的同学。启动 GUI 后,先点「连接」按钮,和服务端建立 socket 连接。默认 token 是 1,输入 1 后点「上传」,客户端就开始训练并上传模型参数。界面上会显示当前轮数、损失值、准确率这些信息。

GUI 底层调用的还是main_client里的训练逻辑,只是把命令行参数变成了按钮和输入框。如果你要改训练参数,还是得去argu.py里改,GUI 只负责触发和展示。

3. 知识蒸馏在客户端怎么落地:教师模型与学生模型的设计

3.1 为什么要用知识蒸馏而不是直接传大模型

联邦学习最怕通信瓶颈。如果每个客户端都传一个 ResNet 级别的模型,一轮下来带宽就炸了。知识蒸馏的思路是:客户端本地有一个大的教师模型和一个小的学生模型,教师模型先训好或者和主任务一起训,然后用教师模型的软标签(soft label)指导学生模型。学生模型参数量可能只有教师模型的十分之一,但精度能接近教师模型。

这套代码里,教师模型和学生模型都定义在model.py里。教师模型层数多、通道宽,学生模型层数少、通道窄。训练时,损失函数是两部分加权:一部分是学生模型输出和真实标签的交叉熵,另一部分是学生模型输出和教师模型输出的 KL 散度。温度参数T控制软标签的平滑程度,T越大,软标签越平滑,学生模型能学到更多类间关系。

3.2 蒸馏损失函数的代码实现

model.py里应该有类似下面的蒸馏损失实现:

import torch import torch.nn as nn import torch.nn.functional as F class DistillationLoss(nn.Module): def __init__(self, temperature=4.0, alpha=0.7): super(DistillationLoss, self).__init__() self.T = temperature self.alpha = alpha self.ce = nn.CrossEntropyLoss() def forward(self, student_out, teacher_out, labels): # 硬标签损失:学生输出和真实标签 hard_loss = self.ce(student_out, labels) # 软标签损失:学生和教师输出的 KL 散度 soft_student = F.log_softmax(student_out / self.T, dim=1) soft_teacher = F.softmax(teacher_out / self.T, dim=1) soft_loss = F.kl_div(soft_student, soft_teacher, reduction='batchmean') * (self.T * self.T) # 加权求和 return self.alpha * soft_loss + (1 - self.alpha) * hard_loss

temperature参数控制软标签的平滑度,常用 3 到 5 之间。alpha控制软硬损失的权重,0.7 表示更偏向教师模型的指导。T * T这个缩放是为了让梯度量级和硬损失匹配,这是 Hinton 那篇蒸馏论文里的标准做法。如果你发现学生模型学不动,先把alpha降到 0.5 试试,让硬标签占主导。

3.3 客户端本地训练循环

客户端训练循环在main_client_1.py里,核心步骤是:从服务端拉取全局模型参数,加载到本地模型,用本地数据训练若干轮,计算蒸馏损失,更新学生模型,最后把学生模型参数上传。下面是一个简化版的训练循环:

for round_idx in range(num_rounds): # 从服务端接收全局模型参数 global_weights = receive_from_server() student_model.load_state_dict(global_weights) # 本地训练 student_model.train() teacher_model.eval() for epoch in range(local_epochs): for batch_x, batch_y in train_loader: optimizer.zero_grad() student_out = student_model(batch_x) with torch.no_grad(): teacher_out = teacher_model(batch_x) loss = distill_criterion(student_out, teacher_out, batch_y) loss.backward() optimizer.step() # 上传学生模型参数 send_to_server(student_model.state_dict())

local_epochs一般设 1 到 3,太大容易过拟合本地数据,导致全局模型发散。teacher_model在本地训练时保持 eval 模式,不更新参数,只提供软标签。如果你想让教师模型也更新,可以在本地先单独训几轮教师模型,再固定住训学生模型。

3.4 模型聚合与参数同步

服务端聚合逻辑在main_server.py里,常见做法是按样本量加权平均:

def aggregate_weights(client_weights_list, sample_counts): total_samples = sum(sample_counts) aggregated = {} for key in client_weights_list[0].keys(): aggregated[key] = sum( client_weights_list[i][key] * (sample_counts[i] / total_samples) for i in range(len(client_weights_list)) ) return aggregated

client_weights_list是各客户端上传的模型参数字典,sample_counts是各客户端的样本数量。加权平均比简单平均更合理,因为样本多的客户端对全局模型的贡献应该更大。聚合完之后,服务端把新参数下发给所有客户端,进入下一轮。

4. 避坑指南:跑这套代码最容易翻车的五个地方

4.1 端口占用导致客户端连不上服务端

现象:客户端启动后一直卡在「等待连接」或者报ConnectionRefusedError

原因:服务端监听的端口被其他程序占用了,或者服务端根本没启动成功。connFun.py里默认端口可能是 9999 或 8888,如果本机有其他服务在用这个端口,就会冲突。

解决:先确认服务端有没有报错,再检查端口占用。Linux 下用lsof -i:9999,Windows 下用netstat -ano | findstr 9999。如果被占用,去argu.pyconnFun.py里把端口改成别的,比如 10086,客户端也要同步改。

4.2 数据路径写死导致换机器就跑不了

现象:换了一台电脑,代码原封不动搬过去,报FileNotFoundError

原因utils.pyinitDate.py里数据路径写的是绝对路径,比如/home/username/FL-IDS-KD/data/KDDTrain+.txt,换机器后路径不存在。

解决:把所有数据路径改成相对路径,基于脚本所在目录拼接。用os.path.dirname(os.path.abspath(__file__))获取当前脚本目录,再拼data子目录。这样不管在哪个机器上跑,只要目录结构不变就能找到数据。

4.3 PyTorch 版本不兼容导致加载权重失败

现象load_state_dictRuntimeError: Error(s) in loading state_dict,或者某些层名对不上。

原因:训练时用的 PyTorch 版本和加载时不一致,或者模型定义改过但权重文件还是旧的。__pycache__里的.pyc是 Python 3.6 编译的,如果你用 3.9 跑,某些语法可能不兼容。

解决:统一用 Python 3.6 到 3.8,PyTorch 用 1.4 到 1.8。如果权重文件对不上,删掉旧的.weight文件重新训。local_model.weightlocal_testModel.weight是训练过程中生成的,不是必须的初始文件。

4.4 客户端数量与代码里写死的不一致

现象:只开了一个客户端,服务端一直等第二个,训练不开始。

原因main_server.py里可能写死了num_clients = 2,必须等够两个客户端连接才进入聚合阶段。

解决:要么开够两个客户端,要么去argu.py里把num_clients改成 1。但联邦学习至少两个客户端才有意义,一个客户端就退化成普通本地训练了。建议还是开两个,用main_client_1.pymain_client_2.py分别跑。

4.5 知识蒸馏温度参数设太大导致学生模型学偏

现象:学生模型准确率一直上不去,甚至比不用蒸馏还差。

原因:温度T设得太大,软标签过于平滑,学生模型学到的类间关系太模糊,反而丢了硬标签的判别信息。

解决:先把T从默认值降到 2 或 3,alpha从 0.7 降到 0.5,让硬标签占更大权重。跑几轮看准确率曲线,如果学生模型和教师模型的差距在缩小,说明蒸馏有效;如果差距扩大,继续降Talpha

5. 进阶玩法:改造成多分类检测与自定义数据集验证

5.1 从二分类扩展到多分类攻击检测

NSL-KDD 的标签不止正常和攻击两类,还有 DoS、Probe、R2L、U2R 四种攻击类型。原始代码可能只做了二分类,但改多分类不难。utils.py里标签映射那部分,把label列从 0/1 改成 0 到 4 的五类。model.py里输出层从 2 个神经元改成 5 个。损失函数用CrossEntropyLoss就行,它自动处理多分类。

改完之后,评估指标也要跟着变。二分类看准确率和 F1 就够了,多分类要看每个类别的召回率,尤其是 U2R 这种样本极少的类别。我一般会打印混淆矩阵,直观看到哪些类别容易混。

5.2 用自定义流量数据替换 NSL-KDD

如果你想用自己的数据跑,格式对齐 NSL-KDD 就行:41 个特征列加 1 个标签列,符号特征做 one-hot,数值特征做归一化。utils.py里的preprocess函数可以复用,只要列名和顺序对上。

替换数据后,客户端样本量可能不均衡,这时候聚合权重按样本量加权就更重要了。如果某个客户端样本特别少,可以在argu.py里给它设一个最低权重,避免它被其他客户端淹没。

5.3 验证蒸馏是否真的起了作用

跑完训练后,resultCompare1.pngresultCompare2.png会生成对比图。我一般会做三组对比:只用联邦学习不用蒸馏、只用蒸馏不用联邦学习、两者都用。看准确率和收敛速度的差异。如果两者都用比单独用效果好,说明蒸馏确实帮学生模型学到了教师模型的泛化能力。

还有一个验证方法是看通信量。学生模型参数量比教师模型少多少,上传下载时间就少多少。在main_client里加个计时,打印每轮上传下载耗时,对比一下用蒸馏和不用蒸馏的通信开销。

从那以后我每次跑联邦学习项目,都会先把num_rounds设成 3 跑一遍全流程,确认数据加载、模型聚合、参数同步都没问题,再加大轮数正式训。这个习惯帮我省了很多等半天才发现报错的时间。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询