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

如何用TensorFlow训练MNIST预训练模型及实现Fashion MNIST自动训练

一、使用TensorFlow训练基于MNIST的预训练模型

MNIST是手写数字数据集,基于预训练模型训练时通常采用迁移学习,复用在通用数据集上训练好的CNN特征提取能力,再适配MNIST的输入输出。以下是具体实现步骤:

  1. 加载并预处理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)
    
  2. 加载预训练基础模型并冻结权重
    以MobileNetV2为例,移除顶部分类层,冻结特征提取部分的权重:

    base_model = tf.keras.applications.MobileNetV2(
        input_shape=(28, 28, 3),
        include_top=False,
        weights='imagenet'
    )
    base_model.trainable = False  # 冻结基础模型,仅训练后续自定义层
    
  3. 构建完整模型
    在预训练模型顶部添加适配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']
    )
    
  4. 训练与微调
    先训练顶部自定义层,之后可解冻部分基础模型层进行微调,提升精度:

    # 初始训练自定义层
    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)
    
  5. 保存模型
    训练完成后保存模型,方便后续调用:

    model.save('mnist_pretrained_model.h5')
    
二、实现Fashion MNIST项目上传预测后自动触发训练

可以通过修改现有项目的预测流程,在完成预测后将用户上传的带标注数据加入训练集,触发增量训练。以下是基于Flask Web项目的实现示例:

  1. 修改预测接口,添加训练触发逻辑
    在图片上传、预测完成后,存储用户数据并调用训练函数:

    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)
    
  2. 关键注意事项

    • 数据标注: 需用户上传图片时提供真实标签,保证训练数据准确性;若无法获取用户标注,可尝试用预测结果作为弱标签,但需注意误差累积。
    • 训练效率: 每次训练采用增量方式,加载已有模型继续训练,避免从头开始;若训练耗时久,可将训练逻辑放入异步任务队列(如Celery),不阻塞接口响应。
    • 数据管理: 定期清理无效或重复数据,防止数据集过大拖慢训练速度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 10:25:25