自定义Soft-attention+Inception ResNet v2模型为何对不同输入输出相同结果?
问题成因
你遇到的全样本预测为同一类的问题属于典型的模型训练塌缩,由多个代码逻辑bug和配置错误共同导致,核心原因如下:
- 数据预处理重复缩放,分布完全错位:你在
ImageDataGenerator中同时调用了Inception ResNet v2自带的preprocess_input和rescale=1./255。前者本身已经完成了归一化逻辑,会将0-255范围的像素值转换为模型预训练适配的[-1,1]区间,额外加1/255缩放会把输入值压缩到[0, 1/255]的极小范围,和预训练权重适配的数据分布完全不符,骨干网络提取的特征完全失效。 - 学习率设置过高,预训练权重被冲毁:你对整个网络(包含预训练的Inception ResNet v2骨干)使用了0.01的Adam学习率,这个数值对于预训练模型微调来说过大,会直接破坏骨干网络上预训练好的有效权重,加上自定义Soft-attention模块是随机初始化的,大学习率会让参数更新快速陷入局部最优,直接塌缩到单类输出。
- 训练步数配置错误,模型未学习完整数据分布:你硬编码
steps_per_epoch=(len(train_df)/10),没有和实际的batch_size对齐,会导致每个epoch喂入模型的数据量、数据分布都和完整训练集不符,模型极容易过拟合到单类小样本上。同时你设置的EarlyStopping patience为70,而总训练epoch才100,相当于早停机制基本失效,模型塌缩后会继续沿着错误方向更新参数直到完全固化。 - 推理阶段预处理缺失,输入分布不匹配:推理时你直接将原始图像reshape后送入模型,没有做和训练阶段一致的预处理,输入数值范围和训练时完全不一致,必然导致预测结果异常。
- 自定义模块与类权重配置不合理:自定义Soft-attention模块没有做权重范围约束,拼接原特征和注意力特征时容易出现数值不稳定,引发梯度消失/爆炸;类权重是手动拍脑袋设置的,没有按照实际类别样本量计算,会进一步引导模型偏向样本占比高或者权重高的类别。
解决方案
按照以下优先级逐一修正即可解决问题:
- 修正数据预处理逻辑,删除重复缩放:
# 去掉rescale参数,仅保留模型自带的预处理函数 datagen=ImageDataGenerator(preprocessing_function=tf.keras.applications.inception_resnet_v2.preprocess_input)
- 调低学习率,采用分层微调策略:不要一开始就解冻整个骨干网络用大学习率训练,先冻结Inception ResNet v2的权重,用1e-4的学习率训练注意力模块和分类头,5-10轮收敛后再解冻骨干网络,将学习率降到1e-5做整体微调。
# 第一阶段训练头部 irv2.trainable = False opt = tf.keras.optimizers.Adam(learning_rate=1e-4) # 第一阶段训练完成后解冻骨干做微调 # irv2.trainable = True # opt = tf.keras.optimizers.Adam(learning_rate=1e-5) model.compile(optimizer=opt, loss='categorical_crossentropy', metrics=['accuracy'])
- 修正训练步数与早停配置,对齐实际batch size:
batch_size = 32 # 替换为你实际使用的批量大小 history = model.fit( train_batches, steps_per_epoch=len(train_df)//batch_size, epochs=30, # 第一阶段训练轮次不要设太高 verbose=1, validation_data=test_batches, validation_steps=len(test_df)//batch_size, callbacks=[checkpoint, EarlyStopping(monitor='val_loss', mode='min',patience=10, min_delta=0.001)], class_weight=class_weights )
- 修正推理阶段逻辑,严格对齐训练预处理:
preprocess = tf.keras.applications.inception_resnet_v2.preprocess_input predictions = [] for i in range(len(arr)): img = arr[i].reshape(1,299,299,3) img = preprocess(img) # 必须做和训练完全一致的预处理 p = model.predict(img, verbose=0) predictions.append(np.argmax(p))
- 基线验证与配置调优:先去掉自定义Soft-attention模块,用全局平均池化接分类头跑通纯Inception ResNet v2基线,确认基线无单类输出问题后再加入注意力模块,注意力分支输出要加sigmoid/softmax激活把权重约束在0-1区间,避免数值不稳定;类权重按照公式
总样本数/(类别数*单类样本数)计算,不要手动赋值。
内容的提问来源于stack exchange,提问作者Javohir
相关产品推荐
相关产品推荐

