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

基于VGG16的MNIST迁移学习内核频繁崩溃问题求助

解决迁移学习MNIST时GPU内核崩溃问题

你遇到的内核崩溃不是硬件问题,完全是代码逻辑的结构错误导致的。下面直接说问题和修复方案:

代码核心错误点

  1. 自定义Model嵌套Sequential导致计算冗余:你在__init__里创建了一个包含base_model和自定义层的Sequential模型,又在call方法里手动遍历这个Sequential的每一层重新计算,等于把所有层的计算重复执行了一遍,会瞬间撑爆内存,直接导致内核崩溃。
  2. call方法参数处理不规范:用args[0]获取输入不符合Keras Model的标准写法,容易引发张量计算异常。
  3. 训练/测试输出不一致:训练时返回logits,测试时返回(logits, prob),这种不一致会让Keras的训练流程混乱,加剧内存问题。

修复后的完整代码

重构自定义Model类

import tensorflow as tf
from tensorflow.keras.applications import VGG16
from tensorflow.keras import layers, models

class VGG16TransferLearning(tf.keras.Model):
    def __init__(self, base_model):
        super(VGG16TransferLearning, self).__init__()
        # 基础模型
        self.base_model = base_model
        # 自定义顶层
        self.flatten = layers.Flatten()
        self.dense1 = layers.Dense(512, activation='relu')
        self.dense2 = layers.Dense(512, activation='relu')
        self.dense3 = layers.Dense(10)  # 输出logits

    def call(self, inputs, training=False):
        # 直接按顺序执行层计算,不需要嵌套Sequential
        out = self.base_model(inputs, training=training)
        out = self.flatten(out)
        out = self.dense1(out)
        out = self.dense2(out)
        logits = self.dense3(out)
        
        # 训练时只返回logits,测试时返回logits和概率(按需调整)
        if training:
            return logits
        else:
            prob = tf.nn.softmax(logits)
            return logits, prob

实例化与训练代码

# MNIST是单通道,已转成3通道输入(75,75,3),注意和VGG16预处理对齐
base_model = VGG16(weights="imagenet", include_top=False, input_shape=(75,75,3))
base_model.trainable = False

# 实例化模型
model = VGG16TransferLearning(base_model)

# 编译时注意:损失函数用from_logits=True,因为我们返回的是logits
model.compile(
    loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    optimizer=tf.keras.optimizers.legacy.Adam(),
    metrics=['accuracy']
)

# 训练
model.fit(x_train, y_train, epochs=10, validation_data=(x_test, y_test))

额外优化建议

  1. VGG16输入预处理:ImageNet的预处理要求像素值缩放到[-1,1]或[0,255]后减均值,建议加上:
    preprocess_input = tf.keras.applications.vgg16.preprocess_input
    x_train = preprocess_input(x_train)
    x_test = preprocess_input(x_test)
    
  2. 输入尺寸调整:VGG16的原始输入是(224,224,3),如果你的输入是(75,75,3),可以在base_model前加一个layers.Resizing(224,224)层,提升特征提取效果:
    class VGG16TransferLearning(tf.keras.Model):
        def __init__(self, base_model):
            super().__init__()
            self.resize = layers.Resizing(224, 224)  # 新增
            self.base_model = base_model
            # ... 其他层不变
        
        def call(self, inputs, training=False):
            out = self.resize(inputs)  # 先调整尺寸
            out = self.base_model(out, training=training)
            # ... 后续计算不变
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 03:39:53