如何用TensorFlow训练MNIST预训练模型及实现Fashion MNIST自动训练
一、使用TensorFlow训练基于MNIST的预训练模型
MNIST是手写数字数据集,基于预训练模型训练时通常采用迁移学习,复用在通用数据集上训练好的CNN特征提取能力,再适配MNIST的输入输出。以下是具体实现步骤:
加载并预处理MNIST数据
将灰度图转换为3通道(适配多数预训练模型的输入要求),同时标准化数据:import tensorflow as tf from tensorflow.keras import layers, models # 加载MNIST数据集 (train_images, train_labels), (test_images, test_labels) = tf.keras.datasets.mnist.load_data() # 预处理:归一化+扩展为3通道 train_images = tf.expand_dims(train_images, axis=-1) train_images = tf.repeat(train_images, 3, axis=-1) train_images = train_images / 255.0 test_images = tf.expand_dims(test_images, axis=-1) test_images = tf.repeat(test_images, 3, axis=-1) test_images = test_images / 255.0 # 标签转为独热编码 train_labels = tf.one_hot(train_labels, depth=10) test_labels = tf.one_hot(test_labels, depth=10)加载预训练基础模型并冻结权重
以MobileNetV2为例,移除顶部分类层,冻结特征提取部分的权重:base_model = tf.keras.applications.MobileNetV2( input_shape=(28, 28, 3), include_top=False, weights='imagenet' ) base_model.trainable = False # 冻结基础模型,仅训练后续自定义层构建完整模型
在预训练模型顶部添加适配MNIST的分类层:model = models.Sequential([ base_model, layers.GlobalAveragePooling2D(), layers.Dense(128, activation='relu'), layers.Dense(10, activation='softmax') ]) model.compile( optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'] )训练与微调
先训练顶部自定义层,之后可解冻部分基础模型层进行微调,提升精度:# 初始训练自定义层 model.fit(train_images, train_labels, epochs=5, batch_size=32, validation_split=0.1) # 微调:解冻基础模型后半部分层 base_model.trainable = True fine_tune_at = 100 # 从第100层开始解冻 for layer in base_model.layers[:fine_tune_at]: layer.trainable = False # 用更小的学习率重新编译,避免破坏已训练的特征 model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5), loss='categorical_crossentropy', metrics=['accuracy'] ) # 继续训练 model.fit(train_images, train_labels, epochs=10, batch_size=32, validation_split=0.1, initial_epoch=5)保存模型
训练完成后保存模型,方便后续调用:model.save('mnist_pretrained_model.h5')
二、实现Fashion MNIST项目上传预测后自动触发训练
可以通过修改现有项目的预测流程,在完成预测后将用户上传的带标注数据加入训练集,触发增量训练。以下是基于Flask Web项目的实现示例:
修改预测接口,添加训练触发逻辑
在图片上传、预测完成后,存储用户数据并调用训练函数:from flask import Flask, request, jsonify import tensorflow as tf import numpy as np from PIL import Image import os app = Flask(__name__) # 加载现有模型 model = tf.keras.models.load_model('fashion_mnist_model.h5') # 存储用户上传的带标注数据目录 user_data_dir = 'user_uploaded_data' os.makedirs(user_data_dir, exist_ok=True) # 预测接口 @app.route('/predict', methods=['POST']) def predict(): if 'image' not in request.files or 'label' not in request.form: return jsonify({'error': '缺少图片或标签参数'}), 400 image_file = request.files['image'] true_label = int(request.form['label']) # 处理图片为模型输入格式 img = Image.open(image_file).convert('L') img = img.resize((28, 28)) img_array = np.array(img) / 255.0 img_array = np.expand_dims(img_array, axis=(0, -1)) # 执行预测 pred_probs = model.predict(img_array) pred_label = np.argmax(pred_probs) # 保存用户上传的带标注数据 img_filename = f'{len(os.listdir(user_data_dir))}_{true_label}.png' img.save(os.path.join(user_data_dir, img_filename)) # 触发自动训练 auto_train() return jsonify({'predicted_label': int(pred_label), 'true_label': true_label}) # 自动增量训练函数 def auto_train(): # 加载原有Fashion MNIST训练数据 (train_images, train_labels), _ = tf.keras.datasets.fashion_mnist.load_data() train_images = train_images / 255.0 train_labels = tf.one_hot(train_labels, depth=10) # 加载用户上传的带标注数据 user_images = [] user_labels = [] for filename in os.listdir(user_data_dir): if filename.endswith('.png'): label = int(filename.split('_')[1].split('.')[0]) img = Image.open(os.path.join(user_data_dir, filename)).convert('L') img_array = np.array(img) / 255.0 user_images.append(img_array) user_labels.append(label) user_images = np.array(user_images) user_labels = tf.one_hot(user_labels, depth=10) # 合并原有数据与用户数据 combined_images = np.concatenate([train_images, user_images], axis=0) combined_labels = np.concatenate([train_labels, user_labels], axis=0) # 增量训练(使用小学习率避免模型遗忘) model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5), loss='categorical_crossentropy', metrics=['accuracy'] ) model.fit(combined_images, combined_labels, epochs=2, batch_size=32, validation_split=0.1) # 保存更新后的模型 model.save('fashion_mnist_model.h5') if __name__ == '__main__': app.run(debug=True)关键注意事项
- 数据标注: 需用户上传图片时提供真实标签,保证训练数据准确性;若无法获取用户标注,可尝试用预测结果作为弱标签,但需注意误差累积。
- 训练效率: 每次训练采用增量方式,加载已有模型继续训练,避免从头开始;若训练耗时久,可将训练逻辑放入异步任务队列(如Celery),不阻塞接口响应。
- 数据管理: 定期清理无效或重复数据,防止数据集过大拖慢训练速度。
内容的提问来源于stack exchange,提问作者Frogleim1
相关产品推荐
相关产品推荐

