如何在Keras中适配数据集训练模型实现网络攻击检测?
网络攻击检测模型拟合问题修复方案
错误1:样本基数模糊、x/y样本数量不一致
该错误由错误的传参逻辑和特征标签拆分逻辑错误共同导致:
model.fit()前两个入参要求为训练集特征矩阵、训练集标签数组,直接传入整份train/test数据集会触发参数匹配错误- 你当前的拆分逻辑存在索引冲突:取0~16列为特征同时取第16列为标签,会导致标签列被重复计入特征,且pandas切片为左闭右开规则,需明确边界:
- 特征列取索引0到16的写法为
df.iloc[:, 0:17],对应17个特征,标签列需改为索引17 - 若标签确实为索引16的列,特征列需改为取0到16的左闭右开切片
df.iloc[:, 0:16]
- 特征列取索引0到16的写法为
错误2:NumPy数组转Tensor失败(int类型不支持,转float仍无效)
按以下优先级排查修复:
- 先做数据清洗:检查数据集是否存在空值、字符串类型的脏数据,强行转float会将空值转为NaN,TensorFlow不允许输入存在非数值
执行以下代码核对数据合法性:
空值可直接删除或用均值/中位数填充,非数值列做编码或删除处理# 检查空值 print(df.isnull().sum()) # 检查列类型,不存在object类型即为合法 print(df.dtypes) - 适配Conv1D输入维度要求:Conv1D要求输入维度为
(样本数, 时间步长, 特征数),你拆分得到的特征为二维数组(样本数, 特征数),缺少最后一个维度,需手动扩展:x_train = x_train.astype('float32')[..., np.newaxis] x_test = x_test.astype('float32')[..., np.newaxis] - 匹配损失函数与标签格式:你当前使用的
categorical_crossentropy要求标签为one-hot编码格式,若attack字段为0/1的整数标签,两种方案二选一:- 损失函数替换为
sparse_categorical_crossentropy,无需修改标签格式 - 对标签做one-hot编码:
from tensorflow.keras.utils import to_categorical y_train = to_categorical(y_train, num_classes=2) y_test = to_categorical(y_test, num_classes=2) - 损失函数替换为
完整可运行参考流程
import pandas as pd import numpy as np from sklearn.model_selection import train_test_split from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Conv1D, MaxPooling1D, Flatten, Dense # 加载+清洗数据 df = pd.read_csv("Bot-IoT数据集路径.csv") df = df.dropna() # 拆分特征标签,此处假设标签为索引16的列 x = df.iloc[:, 0:16].astype('float32').values y = df.iloc[:, 16].astype('int32').values # 拆分训练测试集 x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.2, random_state=42) # 扩展维度适配Conv1D x_train = np.expand_dims(x_train, axis=-1) x_test = np.expand_dims(x_test, axis=-1) # 搭建+编译模型 model = Sequential([ Conv1D(32, kernel_size=3, activation='relu', input_shape=(16, 1)), MaxPooling1D(pool_size=2), Flatten(), Dense(64, activation='relu'), Dense(2, activation='softmax') ]) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) # 模型拟合 model.fit(x_train, y_train, batch_size=32, epochs=10, validation_data=(x_test, y_test))
内容的提问来源于stack exchange,提问作者Abdulaziz Almaawy
相关产品推荐
相关产品推荐

