简介:本资源面向计算机、人工智能、通信工程等专业的在校学生与研究人员,提供一套将联邦学习与知识蒸馏结合用于网络入侵检测的完整Python实现方案,并在NSL-KDD数据集上完成验证。包内共63个文件,以12个py源码、26个pyc编译文件、10个txt说明、3个weight权重文件及csv数据集为主,另含日志、图片与压缩包,整体约26.18MB,目录涵盖服务端、客户端、模型定义、参数配置与GUI界面等模块。运行时可先启动服务端,再开启两个客户端窗口进行联邦训练,通过图形界面连接并上传token即可开始。已有231人学习关注。读者可获得一套可直接运行的入侵检测实验代码、联邦聚合与知识蒸馏的落地思路、NSL-KDD数据处理流程及模型权重,便于复现实验、二次开发或用于课程设计与毕设参考。
1. 联邦学习遇上知识蒸馏:入侵检测模型为什么值得这样搭
单机训练的入侵检测模型,在 NSL-KDD 上刷到 99% 准确率并不难,难的是把它放到真实网络里还能用。真实场景里流量数据分散在多个网段、多个机房,谁都不愿意把原始流量交出去,一是合规,二是带宽,三是数据本身就是资产。联邦学习解决的就是这个「数据不动、模型动」的问题,而知识蒸馏解决的是联邦聚合之后模型太大、推理太慢、边缘设备跑不动的问题。这两个技术叠在一起做网络入侵检测,是这两年被反复讨论的组合,也是很多安全方向毕业设计和横向项目的首选架构。
这篇笔记面向三类人:手里有 NSL-KDD 数据集想跑通一个完整联邦入侵检测流程的;已经会单机训练、想搞清楚联邦聚合和蒸馏到底怎么接的;以及被「灾难性遗忘」「非独立同分布」这些词绕晕、想看到能跑代码的。我会按「先讲清为什么这么选,再给能抄的代码和参数」的节奏走,中间穿插我自己踩过的坑。读完你应该能在一台普通笔记本上,用 Python 把联邦 + 蒸馏的入侵检测流程完整跑一遍,并且知道每个参数动了会发生什么。
2. 数据与任务拆解:NSL-KDD 到底该怎么喂给联邦模型
2.1 NSL-KDD 的字段结构和标签分布
NSL-KDD 是 KDD Cup 99 的修正版,去掉了原始数据里大量重复记录,训练集 KDDTrain+ 大约 12.6 万条,测试集 KDDTest+ 大约 2.2 万条。每条记录 41 个特征加 1 个标签,特征分四类:TCP 连接基本特征(duration、protocol_type、service、flag 等)、流量特征(src_bytes、dst_bytes)、内容特征(num_failed_logins、root_shell 等)、基于时间的统计特征(count、srv_count、serror_rate 等)。标签有 5 大类:Normal、DoS、Probe、R2L、U2R,其中 R2L 和 U2R 样本极少,是典型的类别不平衡。
做联邦之前必须先做单机预处理,因为联邦只是把训练过程拆开,数据清洗逻辑是一样的。常见做法是把 protocol_type、service、flag 三个类别特征做 one-hot,数值特征做标准化,标签做 5 分类或二分类。二分类(Normal vs Attack)更容易跑出好看的数字,多分类更贴近真实检测需求,建议两个都留一份。
import pandas as pd import numpy as np from sklearn.preprocessing import StandardScaler, LabelEncoder # NSL-KDD 列名,官方给的 41 维 + label + difficulty col_names = ["duration","protocol_type","service","flag","src_bytes","dst_bytes", "land","wrong_fragment","urgent","hot","num_failed_logins","logged_in", "num_compromised","root_shell","su_attempted","num_root","num_file_creations", "num_shells","num_access_files","num_outbound_cmds","is_host_login", "is_guest_login","count","srv_count","serror_rate","srv_serror_rate", "rerror_rate","srv_rerror_rate","same_srv_rate","diff_srv_rate", "srv_diff_host_rate","dst_host_count","dst_host_srv_count", "dst_host_same_srv_rate","dst_host_diff_srv_rate","dst_host_same_src_port_rate", "dst_host_srv_diff_host_rate","dst_host_serror_rate","dst_host_srv_serror_rate", "dst_host_rerror_rate","dst_host_srv_rerror_rate","label","difficulty"] train = pd.read_csv("KDDTrain+.txt", names=col_names) test = pd.read_csv("KDDTest+.txt", names=col_names) # 把攻击细类归到 5 大类,方便多分类 attack_map = { "normal":"Normal", "back":"DoS","land":"DoS","neptune":"DoS","pod":"DoS","smurf":"DoS", "teardrop":"DoS","apache2":"DoS","udpstorm":"DoS","processtable":"DoS","worm":"DoS", "ipsweep":"Probe","nmap":"Probe","portsweep":"Probe","satan":"Probe", "mscan":"Probe","saint":"Probe", "ftp_write":"R2L","guess_passwd":"R2L","imap":"R2L","multihop":"R2L", "phf":"R2L","spy":"R2L","warezclient":"R2L","warezmaster":"R2L", "sendmail":"R2L","named":"R2L","snmpgetattack":"R2L","snmpguess":"R2L", "xlock":"R2L","xsnoop":"R2L","httptunnel":"R2L", "buffer_overflow":"U2R","loadmodule":"U2R","perl":"U2R","rootkit":"U2R", "ps":"U2R","sqlattack":"U2R","xterm":"U2R" } train["label"] = train["label"].map(attack_map) test["label"] = test["label"].map(attack_map) # 类别特征 one-hot,数值特征标准化 cat_cols = ["protocol_type","service","flag"] full = pd.concat([train, test], axis=0) full = pd.get_dummies(full, columns=cat_cols) num_cols = [c for c in full.columns if c not in ["label","difficulty"]] scaler = StandardScaler() full[num_cols] = scaler.fit_transform(full[num_cols]) le = LabelEncoder() full["label_id"] = le.fit_transform(full["label"])这段代码的关键点有三个。第一,one-hot 必须在 train 和 test 合并之后做,否则两边类别不一致会导致列对不齐,这是新手最常翻车的地方。第二,difficulty 列要丢掉,它是官方给的难度标记,不是特征,留着会泄漏。第三,标准化用全量数据 fit 在严格意义上有一点信息泄漏,但在联邦场景下我们本来就要模拟「全局分布」,这样处理可以接受;如果追求严谨,改成只用训练集 fit、测试集 transform。
2.2 把集中式数据切成联邦客户端
联邦学习要模拟多个客户端各自持有数据。NSL-KDD 没有天然的客户端划分,常见做法是按标签做 Non-IID 切分,让每个客户端只拿到部分类别的样本,这样才能真实暴露联邦聚合在非独立同分布下的问题。我一般会切 5 到 10 个客户端,每个客户端数据量在几千到两万之间。
from sklearn.model_selection import train_test_split def split_federated(df, n_clients=5, mode="iid", seed=42): rng = np.random.default_rng(seed) clients = [] if mode == "iid": # 随机均分,每个客户端分布接近全局 idx = rng.permutation(len(df)) chunks = np.array_split(idx, n_clients) for c in chunks: clients.append(df.iloc[c].reset_index(drop=True)) else: # Non-IID:按标签排序后分片,每个客户端只覆盖部分类别 df_sorted = df.sort_values("label_id").reset_index(drop=True) chunks = np.array_split(np.arange(len(df_sorted)), n_clients) for c in chunks: clients.append(df_sorted.iloc[c].reset_index(drop=True)) return clients clients = split_federated(full, n_clients=5, mode="noniid") for i, c in enumerate(clients): print(f"client {i}: {len(c)} samples, label dist={c['label_id'].value_counts().to_dict()}")参数说明:n_clients 控制客户端数量,越多越接近真实跨机构场景,但通信轮次和聚合开销也越大;mode 选 iid 用来验证流程正确性,选 noniid 用来观察联邦的真实难点。我建议先用 iid 跑通,再切 noniid 对比准确率掉多少,这个差值就是你这套方案要优化的空间。
3. 联邦训练主循环:从本地模型到全局聚合
3.1 本地模型结构怎么定
入侵检测是结构化表格数据,不需要上 ResNet 那种大模型。我一般用 3 层全连接:输入维度等于特征数,中间两层 128 和 64,输出 5 类,每层后面接 ReLU 和 Dropout(0.3)。这个规模在 NSL-KDD 上单机就能到 98% 以上,联邦场景下也不会因为模型太大导致聚合后收敛慢。
import torch import torch.nn as nn class IDSModel(nn.Module): def __init__(self, in_dim, n_classes=5): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, 64), nn.ReLU(), nn.Dropout(0.3), nn.Linear(64, n_classes) ) def forward(self, x): return self.net(x)in_dim 由预处理后的特征数决定,one-hot 之后通常在 120 左右。Dropout 放在中间层是为了防止联邦聚合时客户端过拟合本地小数据,这一点在 Non-IID 下尤其明显。
3.2 FedAvg 聚合的最小实现
联邦平均(FedAvg)的核心就一句话:各客户端用本地数据训练若干轮,把模型参数上传,服务端按样本量加权平均。下面是一个不依赖任何联邦框架、纯 PyTorch 手写的版本,方便你理解每一步在干什么。
import copy from torch.utils.data import DataLoader, TensorDataset def local_train(model, df, epochs=1, lr=1e-3, batch_size=256): model = copy.deepcopy(model) model.train() X = torch.tensor(df.drop(columns=["label","label_id"]).values, dtype=torch.float32) y = torch.tensor(df["label_id"].values, dtype=torch.long) loader = DataLoader(TensorDataset(X, y), batch_size=batch_size, shuffle=True) opt = torch.optim.Adam(model.parameters(), lr=lr) loss_fn = nn.CrossEntropyLoss() for _ in range(epochs): for xb, yb in loader: opt.zero_grad() loss = loss_fn(model(xb), yb) loss.backward() opt.step() return model.state_dict(), len(df) def fed_avg(global_state, client_states, client_sizes): total = sum(client_sizes) new_state = copy.deepcopy(global_state) for k in global_state.keys(): new_state[k] = sum( client_states[i][k] * (client_sizes[i] / total) for i in range(len(client_states)) ) return new_state # 主循环 in_dim = clients[0].drop(columns=["label","label_id"]).shape[1] global_model = IDSModel(in_dim) global_state = global_model.state_dict() for rnd in range(20): client_states, client_sizes = [], [] for df_c in clients: global_model.load_state_dict(global_state) st, sz = local_train(global_model, df_c, epochs=1, lr=1e-3) client_states.append(st) client_sizes.append(sz) global_state = fed_avg(global_state, client_states, client_sizes) print(f"round {rnd} done")逻辑说明:每一轮先把全局参数下发到各客户端,客户端本地跑 1 个 epoch,回传 state_dict 和样本数,服务端按样本数加权平均。参数上,本地 epoch 数不要设太大,1 到 3 之间比较稳,设大了客户端会各自跑偏,聚合后反而震荡;学习率 1e-3 是 Adam 的常规起点,Non-IID 严重时可以降到 5e-4。通信轮次 20 轮在 NSL-KDD 上基本能收敛,想省时间可以设 10 轮看趋势。
3.3 聚合权重和客户端采样
真实联邦里不是每轮所有客户端都在线,所以还要加客户端采样。常见做法是每轮随机抽 50% 到 80% 的客户端参与。加权平均的权重除了样本量,也可以按类别均衡度调整,比如某个客户端 R2L 样本多,就给它更高权重,这样能缓解 Non-IID 下少数类被淹没的问题。这一步没有标准答案,属于调参空间,建议先跑通等权版本,再试加权版本对比。
4. 知识蒸馏接进来:让全局模型变小、让客户端变强
4.1 蒸馏在联邦里的两种用法
知识蒸馏在联邦入侵检测里有两种典型接法。第一种是服务端蒸馏:把联邦聚合出来的大模型当教师,蒸馏出一个更小的学生模型部署到边缘设备,解决推理速度问题。第二种是客户端蒸馏:用全局模型当教师,指导每个客户端的本地模型训练,缓解 Non-IID 下本地模型只认自己那几类样本的问题。这两种可以同时用,我一般先做客户端蒸馏,因为它直接改善联邦的收敛质量。
蒸馏的损失函数是软标签损失加硬标签损失的加权和:
def distill_loss(student_logits, teacher_logits, hard_labels, T=4.0, alpha=0.7): # 软标签:教师和学生都过温度 T 的 softmax soft = nn.KLDivLoss(reduction="batchmean")( nn.functional.log_softmax(student_logits / T, dim=1), nn.functional.softmax(teacher_logits / T, dim=1) ) * (T * T) # 硬标签:正常交叉熵 hard = nn.functional.cross_entropy(student_logits, hard_labels) return alpha * soft + (1 - alpha) * hard参数说明:T 是温度,控制软标签的平滑程度,常用 3 到 5,T 越大教师输出的类别间关系越明显,但太大也会引入噪声;alpha 是软标签权重,0.7 意味着更依赖教师,Non-IID 严重时可以把 alpha 提到 0.8,让本地模型更多地向全局知识靠拢。T*T 这个缩放是为了让软标签损失的梯度量级和硬标签匹配,不加的话蒸馏几乎不起作用,这是很多人第一次写蒸馏会漏掉的细节。
4.2 客户端蒸馏的完整训练步骤
把蒸馏接进本地训练,就是把原来的交叉熵换成蒸馏损失,教师模型用上一轮的全局模型。
def local_train_distill(student, teacher, df, epochs=1, lr=1e-3, batch_size=256, T=4.0, alpha=0.7): student = copy.deepcopy(student) teacher.eval() student.train() X = torch.tensor(df.drop(columns=["label","label_id"]).values, dtype=torch.float32) y = torch.tensor(df["label_id"].values, dtype=torch.long) loader = DataLoader(TensorDataset(X, y), batch_size=batch_size, shuffle=True) opt = torch.optim.Adam(student.parameters(), lr=lr) for _ in range(epochs): for xb, yb in loader: opt.zero_grad() with torch.no_grad(): t_logits = teacher(xb) s_logits = student(xb) loss = distill_loss(s_logits, t_logits, yb, T, alpha) loss.backward() opt.step() return student.state_dict(), len(df)逻辑说明:教师模型只做前向、不更新梯度,学生模型同时学软标签和硬标签。注意教师是上一轮的全局模型,不是当前轮聚合后的模型,否则会有信息泄漏。每一轮结束后,学生模型参数上传做 FedAvg,得到新的全局模型,下一轮再当教师。这个循环跑 15 到 20 轮,Non-IID 场景下准确率通常比纯 FedAvg 高 2 到 5 个百分点。
4.3 服务端蒸馏出小模型
如果最终要部署到边缘,还要在服务端做一次蒸馏。教师是联邦聚合后的全局模型,学生是一个更窄的网络,比如 64-32 两层。用全体客户端数据不现实,常见做法是用服务端持有的少量公开数据,或者让客户端上传一批无标签样本的软标签。这里给一个用测试集当代理数据的版本,仅用于验证流程。
def server_distill(teacher, student, proxy_df, epochs=10, lr=1e-3, T=4.0): teacher.eval() student.train() X = torch.tensor(proxy_df.drop(columns=["label","label_id"]).values, dtype=torch.float32) y = torch.tensor(proxy_df["label_id"].values, dtype=torch.long) loader = DataLoader(TensorDataset(X, y), batch_size=256, shuffle=True) opt = torch.optim.Adam(student.parameters(), lr=lr) for _ in range(epochs): for xb, yb in loader: opt.zero_grad() with torch.no_grad(): t_logits = teacher(xb) s_logits = student(xb) loss = distill_loss(s_logits, t_logits, yb, T, alpha=0.9) loss.backward() opt.step() return student服务端蒸馏 alpha 设 0.9,因为代理数据可能和真实分布有偏差,硬标签不可靠,更多依赖教师的软输出。学生模型参数量通常能压到教师的 1/4 到 1/3,推理延迟明显下降,准确率掉 1 个百分点以内是可以接受的。
5. 避坑与排查:联邦蒸馏入侵检测最容易翻车的地方
5.1 客户端准确率很高但全局模型不涨
现象:每个客户端本地训练完准确率都 95% 以上,聚合后全局模型在测试集上只有 80% 出头,甚至比单机还低。原因:Non-IID 切分下各客户端只见过部分类别,本地模型对没见过的类别输出是随机的,加权平均后互相抵消。解决:一是改用客户端蒸馏,用全局模型当教师约束本地更新方向;二是聚合时按类别均衡度加权,让覆盖少数类的客户端权重更高;三是增加客户端采样比例,别每轮只抽 30%。
5.2 蒸馏损失不下降,学生和教师输出几乎一样
现象:加了蒸馏之后 loss 曲线平的,学生模型准确率和不用蒸馏差不多。原因:多半是忘了 TT 缩放,或者 alpha 设得太小,软标签的梯度被硬标签淹没。解决:先确认 distill_loss 里乘了 TT,再把 alpha 从 0.5 往上调到 0.7 到 0.8,观察 loss 是否开始下降。另外教师模型要设 eval 模式,否则 Dropout 会让软标签抖动。
5.3 测试集准确率远低于训练集
现象:训练集 99%,KDDTest+ 上只有 75% 到 80%。原因:NSL-KDD 的测试集里包含训练集没出现过的攻击细类,这是数据集本身的设计,不是模型 bug。解决:接受这个差距,重点看测试集上的召回率和混淆矩阵,尤其是 R2L 和 U2R 的召回。如果要做对比实验,所有方法都用同一个测试集,差距才有意义。别为了刷高测试集准确率去调参,那是自欺欺人。
5.4 通信轮次增加但准确率震荡
现象:联邦训练到第 10 轮之后准确率上下跳,不收敛。原因:本地 epoch 太多或者学习率太大,客户端每轮跑偏太远,聚合后互相拉扯。解决:本地 epoch 降到 1,学习率降到 5e-4,或者加学习率衰减。另一个可能是客户端采样每轮变化太大,固定采样集合或者用更大的采样比例能缓解。
5.5 显存或内存爆掉
现象:客户端数量一多,或者 batch_size 设大,跑几轮就 OOM。原因:每个客户端都 copy 了一份模型,加上 DataLoader 的缓存,内存是客户端数的倍数。解决:客户端串行训练,不要并行;batch_size 从 256 降到 128;数据用 float32 而不是 float64,预处理阶段就转好。笔记本上跑 5 个客户端、12 万条数据,内存占用大概 2 到 3 GB,属于可控范围。
6. 进阶技巧:用灾难性遗忘的视角调联邦蒸馏
「灾难性遗忘 联邦学习」是这两年被提得很多的组合词,放到入侵检测里特别贴切:联邦每聚合一次,全局模型就可能把上一轮学到的某些攻击类型忘掉,尤其是 R2L、U2R 这种样本少的类。我自己的习惯是,每轮聚合后不只看总准确率,而是单独记录每个类别的召回率,画一条按轮次变化的曲线。如果某一类的召回在某轮之后突然掉下去,那就是遗忘发生了。
应对遗忘,我一般用两个手段叠加。第一个是蒸馏时提高 alpha,让本地模型更多保留全局知识,减少被本地数据带偏。第二个是在聚合时做参数层面的约束,把当前轮全局参数和上一轮全局参数做加权,类似 FedProx 的思路:
def fed_avg_with_momentum(prev_state, new_state, mu=0.2): # mu 越大,越保留上一轮知识,缓解遗忘 out = {} for k in new_state.keys(): out[k] = (1 - mu) * new_state[k] + mu * prev_state[k] return outmu 取 0.1 到 0.3 之间,太大模型学不动新东西,太小起不到抗遗忘作用。我通常从 0.2 起步,看各类召回曲线再微调。
验证这套方案有没有效果,别只看一个准确率数字。我习惯做三组对照:纯 FedAvg、FedAvg + 客户端蒸馏、FedAvg + 客户端蒸馏 + 动量聚合,在同一个 Non-IID 切分和同一个测试集上跑,记录总体准确率、宏平均 F1、以及 R2L/U2R 的召回。三组跑下来,如果蒸馏和动量确实把少数类召回拉起来了,这套方案就值得继续投入;如果没差别,说明你的 Non-IID 程度还不够,把客户端数量加到 10 个、切分再偏一点,问题才会暴露出来。
最后说个我自己的教训:一开始我图省事,把预处理后的全量数据直接按行随机切给客户端,跑出来效果特别好,还以为是方案牛。后来才发现那是 IID 切分,等于没模拟联邦的真实难点。换成按标签切之后,准确率掉了快 10 个点,才真正开始理解蒸馏和抗遗忘为什么必要。做联邦方向,切分方式比模型结构重要得多,这一步别偷懒。希望帮到你。
本文还有配套的精品资源,点击获取