基于深度学习的毒蘑菇检测

文章目录

任务介绍

本次任务为毒蘑菇的二元分类,任务本身并不复杂,适合初学者,主要亮点在于对字符数据的处理,还有尝试了加深神经网络深度的效果,之后读者也可自行改变观察效果,比赛路径将于附录中给出。

数据概览

本次任务的数据集比较简单

  • train.csv 训练文件
  • test.csv 测试文件
  • sample_submission.csv 提交示例文件

具体内容就是关于毒蘑菇的各种特征,可在附录中获取数据集。

数据处理

数据读取与拼接

这段代码提取了数据文件,并且对两个不同来源的数据集进行了拼接,当我们的数据集较小时,就可采用这种方法,获取其他的数据集并将两个数据集合并起来。

python 复制代码
import pandas as pd
from sklearn.preprocessing import LabelEncoder, StandardScaler
python 复制代码
file = pd.read_csv("/kaggle/input/playground-series-s4e8/train.csv", index_col="id")
file2 = pd.read_csv("/kaggle/input/mushroom-classification-edible-or-poisonous/mushroom.csv")
python 复制代码
file_all = pd.concat([file, file2])
字符数据转化

这段代码主要就是提取出字符数据,因为字符是无法直接被计算机处理,所以我们提取出来后,再将字符数据映射为数字数据。

python 复制代码
char_features = ['cap-shape', 'cap-surface', 'cap-color', 'does-bruise-or-bleed', 'gill-attachment', 'gill-spacing', 'gill-color', 'stem-root', 'stem-surface', 'stem-color', 'veil-type', 'veil-color', 'has-ring', 'ring-type', 'spore-print-color', 'habitat', 'season']
python 复制代码
for i in char_features:
    file_all[i] = LabelEncoder().fit_transform(file_all[i])
python 复制代码
file_all = file_all.fillna(0)
python 复制代码
train_col = ['cap-diameter', 'stem-height', 'stem-width', 'cap-shape', 'cap-surface', 'cap-color', 'does-bruise-or-bleed', 'gill-attachment', 'gill-spacing', 'gill-color', 'stem-root', 'stem-surface', 'stem-color', 'veil-type', 'veil-color', 'has-ring', 'ring-type', 'spore-print-color', 'habitat', 'season']
python 复制代码
X = file_all[train_col]
y = file_all['class']
标签数据映射

除了用上述方法进行字符转化外,还可以使用map函数,以下是具体操作。

python 复制代码
y.unique()
python 复制代码
# 构建映射字典
applying = {'e': 0, 'p': 1}
python 复制代码
y = y.map(applying)
数据集划分

这段代码使用sklearn库将数据集划分为训练集和测试集。

python 复制代码
from sklearn.model_selection import train_test_split
python 复制代码
x_train, x_test, y_train, y_test = train_test_split(X, y, test_size=0.1)
python 复制代码
x_train.shape, y_train.shape
数据标准化

这段代码将我们的数据进行归一化,减小数字大小方便计算,但是仍然保持他们之间的线性关系,不会对结果产生影响。

python 复制代码
scaler = StandardScaler()
python 复制代码
x_train = scaler.fit_transform(x_train)
x_test = scaler.fit_transform(x_test)

模型构建与训练

这段代码使用torch库构建了深度学习模型,主要运用了线性层,还进行了正则化操作,防止模型过拟合。

模型构建
python 复制代码
import torch
import torch.nn as nn
python 复制代码
class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(20, 256)
        self.relu = nn.ReLU()
        self.dropout = nn.Dropout(p=0.2)
        self.linear1 = nn.Linear(256, 128)
        self.linear2 = nn.Linear(128, 64)
        self.linear3 = nn.Linear(64, 48)
        self.linear4 = nn.Linear(48, 32)
        self.linear5 = nn.Linear(32, 2)
        
    def forward(self, x):
        out = self.linear(x)
        out = self.relu(out)
        out = self.linear1(out)
        out = self.relu(out)
        out = self.dropout(out)
        out = self.linear2(out)
        out = self.relu(out)
        out = self.linear3(out)
        out = self.dropout(out)
        out = self.relu(out)
        out = self.linear4(out)
        out = self.relu(out)
        out = self.linear5(out)
        
        return out

对模型类进行实例化。

python 复制代码
model = Model()
数据批处理

由于数据一条一条的处理起来很慢,因此我们可以将数据打包,一次给模型输入多条数据,能有效节省时间。

python 复制代码
import torch.nn.functional as F
python 复制代码
class Dataset(torch.utils.data.Dataset):
    def __init__(self, x, y):
        self.x = x
        self.y = y
        
    def __len__(self):
        return len(self.x)
    
    def __getitem__(self, i):
        x = torch.Tensor(self.x[i])
        y = torch.tensor(self.y.iloc[i])
        
        return x, y
python 复制代码
train_data = Dataset(x_train, y_train)
test_data = Dataset(x_test, y_test)
python 复制代码
loader = torch.utils.data.DataLoader(train_data, batch_size=64, drop_last=True, shuffle=True)
python 复制代码
test_loader = torch.utils.data.DataLoader(test_data, batch_size=256, drop_last=True, shuffle=True)
模型训练

这段代码就是模型的训练过程,包括创建优化器,定义损失函数等,还在训练过程中测试准确率与损失函数值,动态的观察训练过程。

python 复制代码
from tqdm import tqdm
import matplotlib.pyplot as plt
python 复制代码
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.1, weight_decay=1e-5)
python 复制代码
from sklearn.metrics import matthews_corrcoef
python 复制代码
flag = 0
for i in range(10):
    for x, label in tqdm(loader):
        out = model(x)
        loss = criterion(out, label)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()
        flag+=1
        if flag%500 == 0:
            test = next(iter(test_loader))
            t_out = model(test[0]).argmax(dim=1)
            print("loss=", loss.item())
            acc = (t_out == test[1]).sum().item()/len(test[1])
            mcc = matthews_corrcoef(t_out, test[1])
            print("acc=", acc)
            print("mcc=", mcc)
        

文件提交

这段代码主要就是使用训练好的模型在测试集上预测,并且将其整合成提交文件。

python 复制代码
test_file = pd.read_csv("/kaggle/input/playground-series-s4e8/test.csv")
python 复制代码
for i in char_features:
    test_file[i] = LabelEncoder().fit_transform(test_file[i])
python 复制代码
test_file.fillna(0)
python 复制代码
test_x = torch.Tensor(test_file[train_col].values)
test_x = torch.Tensor(scaler.fit_transform(test_x))
python 复制代码
out = model(test_x)
out = pd.Series(out.argmax(dim=1))
python 复制代码
map2 = {0: 'e', 1: 'p'}
python 复制代码
result = out.map(map2)
python 复制代码
answer = pd.DataFrame({'id': test_file['id'], "class": result})
python 复制代码
answer.to_csv('submission.csv', index=False)

结果

将文件提交后,得到了0.97的成绩,已经非常接近1了,证明模型的效果非常不错。

附录

比赛链接:https://www.kaggle.com/competitions/playground-series-s4e8

额外数据集地址:https://www.kaggle.com/datasets/vishalpnaik/mushroom-classification-edible-or-poisonous

相关推荐
JJJennie77721 分钟前
当AI开始学会欺骗
网络·人工智能
糖果店的幽灵30 分钟前
大模型测评DeepEval快速入门-安全与通用指标详解
人工智能·安全·langgraph·大模型测评·deepeval
江畔柳前堤6 小时前
roLabelImg 详细安装教程
开发语言·人工智能·后端·云原生
阿里云大数据AI技术6 小时前
分链路差异化设计的DSP准实时数仓|钛动科技基于阿里云实时计算 Flink 版 + DLF Paimon + EMR Serverless StarRocks 的实践
人工智能·flink
陕西企来客6 小时前
2026年7月AI智能搜索曝光趋势研判
大数据·人工智能·机器学习·ai智能搜索曝光
阿里云大数据AI技术7 小时前
从算力到智能体,面向 Agentic AI 的基础设施演进
人工智能·agent
hangyuekejiGEO8 小时前
GEO技术服务选型指南
大数据·人工智能·python
阿里云大数据AI技术8 小时前
EMR Serverless Spark AI Function 的双维降本实践
人工智能·sql·spark
维基框架8 小时前
GitHub源码处理提速 一趟扫描反而更慢
人工智能·github
冬奇Lab8 小时前
代码库知识库系列(05):向量检索 vs 知识图谱——加了调用图并没有变更好
人工智能