You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Keras的CNN猫狗图像分类器中新增青蛙类别?

将Keras二分类CNN扩展为三分类(猫、狗、青蛙)的步骤

没问题,把原本的猫/狗二分类CNN改成猫/狗/青蛙三分类其实只需要调整几个关键部分,我给你一步步说清楚:

第一步:调整数据集结构

首先你需要整理你的训练集和测试集文件夹,确保它们的结构是这样的:

training_set/
├── cats/
│   ├── cat001.jpg
│   └── ...
├── dogs/
│   ├── dog001.jpg
│   └── ...
└── frogs/
    ├── frog001.jpg
    └── ...

test_set/
├── cats/
├── dogs/
└── frogs/

每个子文件夹里放对应类别的图片,Keras的flow_from_directory会自动根据子文件夹名称识别类别。

第二步:修改模型的输出层

原来的二分类用的是单神经元+sigmoid激活,三分类需要改成3个神经元+softmax激活(softmax会输出每个类别的概率,总和为1):

# 替换原来的输出层
classifier.add(Dense(units = 3, activation = 'softmax'))

第三步:调整损失函数和数据生成器的类别模式

  • 损失函数从二分类的binary_crossentropy换成多分类的categorical_crossentropy
  • 数据生成器的class_mode从'binary'改成'categorical',这样生成器会输出one-hot编码的标签,和损失函数匹配

第四步:修改ModelCheckpoint的监控指标(可选)

如果你的Keras版本较新,val_acc可能已经被更名为val_accuracy,如果训练时出现指标找不到的错误,记得把monitor='val_acc'改成monitor='val_accuracy'。


修改后的完整代码

from keras.models import Sequential
from keras.layers import Conv2D
from keras.layers import MaxPooling2D
from keras.layers import Flatten
from keras.layers import Dense
from keras.callbacks import ModelCheckpoint
from keras.preprocessing.image import ImageDataGenerator

# 搭建CNN模型
classifier = Sequential()
classifier.add(Conv2D(32, (3, 3), input_shape = (64, 64, 3), activation = 'relu'))
classifier.add(MaxPooling2D(pool_size = (2, 2)))
classifier.add(Conv2D(32, (3, 3), activation = 'relu'))
classifier.add(MaxPooling2D(pool_size = (2, 2)))
classifier.add(Flatten())
classifier.add(Dense(units = 128, activation = 'relu'))
# 关键修改1:输出层改为3个神经元+softmax
classifier.add(Dense(units = 3, activation = 'softmax'))

# 关键修改2:损失函数换成categorical_crossentropy
classifier.compile(optimizer = 'adam', loss = 'categorical_crossentropy', metrics = ['accuracy'])

# 模型保存回调(如果新版本Keras报错,把val_acc改成val_accuracy)
filepath="weights-improvement-{epoch:02d}-{val_acc:.2f}.hdf5"
checkpoint=ModelCheckpoint(filepath, monitor='val_acc', verbose=1, save_best_only=True, mode='max')
callback_list=[checkpoint]

# 数据增强
train_datagen = ImageDataGenerator(rescale = 1./255, shear_range = 0.2, zoom_range = 0.2, horizontal_flip = True)
test_datagen = ImageDataGenerator(rescale = 1./255)

# 关键修改3:class_mode改为categorical
training_set = train_datagen.flow_from_directory('training_set', target_size = (64, 64), batch_size = 32, class_mode = 'categorical')
test_set = test_datagen.flow_from_directory('test_set', target_size = (64, 64), batch_size = 32, class_mode = 'categorical')

# 训练模型
classifier.fit_generator(training_set, steps_per_epoch = 8000, epochs = 10, validation_data = test_set, validation_steps = 2000, callbacks=callback_list)

# 保存训练好的模型
classifier.save('model_after_trained.h5')

额外提示

  • 训练前可以打印training_set.class_indices,确认Keras识别的类别对应关系(比如{'cats':0, 'dogs':1, 'frogs':2}),方便后续预测时对应结果
  • 如果你的数据集大小变化了,记得调整steps_per_epoch和validation_steps的数值,一般是总样本数除以batch_size

内容的提问来源于stack exchange,提问作者Arjun

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 10:00:36