MNIST数据集:10个二分类神经网络实现的代码优化建议咨询
嘿,这个思路挺有意思的——用10个二分类器实现多分类,刚好适合初学者拆解多分类任务的本质!我来给你一步步拆解怎么修改代码实现这个方案:
第一步:MNIST数据集的类别分离处理
首先要把原MNIST的多分类标签,转换成10组二分类标签(每组对应“当前类别vs其余所有类别”)。这里可以用numpy的where函数快速实现:
import numpy as np from tensorflow.keras.datasets import mnist # 加载并预处理基础数据 (x_train, y_train), (x_test, y_test) = mnist.load_data() # 归一化+展平28*28图像为一维向量 x_train = x_train.astype('float32') / 255.0 x_test = x_test.astype('float32') / 255.0 x_train = x_train.reshape((len(x_train), 28*28)) x_test = x_test.reshape((len(x_test), 28*28)) # 生成10组二分类数据集 binary_datasets = [] for target_class in range(10): # 训练集标签:目标类别标记为1,其余为0 y_train_binary = np.where(y_train == target_class, 1, 0) # 测试集标签同理 y_test_binary = np.where(y_test == target_class, 1, 0) binary_datasets.append( (x_train, y_train_binary, x_test, y_test_binary) )
第二步:构建通用的二分类神经网络
因为10个模型的结构可以完全一致,我们写一个函数来生成复用的二分类模型,避免重复代码:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense, Dropout def build_binary_classifier(input_shape=(784,)): model = Sequential([ Dense(128, activation='relu', input_shape=input_shape), Dropout(0.2), # 防止过拟合 Dense(64, activation='relu'), Dropout(0.2), Dense(1, activation='sigmoid') # 二分类用sigmoid输出概率 ]) # 二分类专用的损失函数和优化器 model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) return model
第三步:循环训练10个二分类模型
接下来遍历每组二分类数据集,训练对应的模型并保存:
trained_models = [] # 存储10个训练好的模型 for idx, (x_tr, y_tr, x_te, y_te) in enumerate(binary_datasets): print(f"===== 训练第 {idx} 个模型(识别数字 {idx})=====") model = build_binary_classifier() # 加入早停回调防止过拟合(可选但推荐) from tensorflow.keras.callbacks import EarlyStopping early_stop = EarlyStopping(monitor='val_loss', patience=3, restore_best_weights=True) model.fit(x_tr, y_tr, epochs=15, batch_size=32, validation_data=(x_te, y_te), callbacks=[early_stop], verbose=1) trained_models.append(model)
第四步:集成模型完成多分类预测
训练完成后,我们可以通过“取最大概率”的方式集成10个二分类器的结果:每个模型输出样本属于对应类别的概率,最终预测概率最高的类别。
def predict_digit(x_sample, models): # 收集每个模型的预测概率 prob_list = [] for model in models: # 预测单个样本的概率 prob = model.predict(x_sample.reshape(1, 784), verbose=0)[0][0] prob_list.append(prob) # 取概率最大的索引作为最终类别 return np.argmax(prob_list) # 测试一个样本 test_idx = 100 true_label = y_test[test_idx] pred_label = predict_digit(x_test[test_idx], trained_models) print(f"真实数字: {true_label}, 预测数字: {pred_label}")
额外优化建议
- 如果想简化代码,可以用
sklearn的OneVsRestClassifier封装Keras模型,但手动实现10个模型更适合理解原理 - 可以给不同模型调整超参数(比如神经元数量、dropout比例),观察对分类效果的影响
- 可以把训练好的模型保存下来,后续直接加载使用:
model.save(f"digit_{idx}_classifier.h5")
内容的提问来源于stack exchange,提问作者7speed
相关产品推荐
相关产品推荐

