基于ML的四边形角点检测问题求助:模型训练异常排查
四边形角点检测(含缺失角)问题解决方案
核心问题分析
当前方案的主要问题在于:
- 直接用全连接层回归坐标丢失了空间上下文信息,角点的位置依赖图像局部空间特征,
Flatten操作会破坏这种关联 - 损失函数选择不匹配归一化后的坐标,MSLE更适合大数值范围的回归,0-1区间的坐标用MSE/Huber Loss更直接
- 模型特征提取能力不足,浅卷积+小全连接层无法捕捉折角这类局部形变的精细特征
具体改进方案
1. 改用热图回归方案(推荐)
这是工业界关键点检测的主流方案,能天然处理角点缺失的情况:
- 将每个角点标注转换为高斯热图:在角点坐标处生成一个高斯分布的热图,缺失角点的热图全为0
- 模型输出4个与输入下采样后尺寸匹配的热图(比如224x224输入下采样到56x56,输出4个56x56热图)
- 推理时,对每个热图取峰值坐标,再通过下采样倍数映射回原图尺寸;热图峰值低于阈值则判定该角点缺失
2. 优化直接回归模型的结构
如果坚持直接回归坐标:
- 用
GlobalAveragePooling2D替代Flatten,保留空间特征的同时减少参数 - 复用预训练骨干(如MobileNetV3Small),去掉分类头后接回归层,利用预训练的通用图像特征
- 增大全连接层的神经元数量(比如从20改为256/128),避免特征过度压缩
3. 调整损失函数
- 坐标归一化到0-1区间后,优先使用MSE(均方误差)或Huber Loss:Huber Loss对训练初期的离群预测更鲁棒
- 若使用热图方案,用MSE计算预测热图与高斯标注热图的差异;或用Focal Loss缓解热图中背景区域(非角点)占比过高的问题
4. 强化数据增强
针对纸张折角场景添加针对性增强:
- 随机旋转、平移、缩放、轻微透视变换(模拟纸张摆放角度)
- 随机调整亮度、对比度、添加噪声(模拟不同拍摄环境)
- 增加缺失角点样本的训练比例,让模型学习区分角点存在/缺失的特征
5. 优化输出层设计
直接回归时,可在输出层增加角点置信度:
- 总输出设为12维:4个(x,y)坐标 + 4个置信度(0表示缺失,1表示存在)
- 损失函数采用组合损失:坐标回归损失(MSE) + 置信度分类损失(交叉熵),权重可根据需求调整
改进后模型示例(热图方案)
from tensorflow.keras import layers, models, losses # 输入尺寸:224x224x3 inputs = layers.Input(shape=(224, 224, 3)) # 复用预训练MobileNetV3Small作为特征提取骨干 base_model = models.MobileNetV3Small( input_shape=(224, 224, 3), include_top=False, weights='imagenet' ) # 冻结骨干前几层,只训练顶部特征层 for layer in base_model.layers[:-5]: layer.trainable = False x = base_model(inputs) # 上采样恢复空间尺寸,输出4个角点的热图 x = layers.Conv2DTranspose(64, (3,3), strides=(2,2), padding='same', activation='relu')(x) x = layers.Conv2DTranspose(32, (3,3), strides=(2,2), padding='same', activation='relu')(x) # 用sigmoid确保热图输出在0-1区间 outputs = layers.Conv2D(4, (1,1), activation='sigmoid')(x) model = models.Model(inputs, outputs) model.compile(optimizer='adam', loss=losses.MeanSquaredError())
内容的提问来源于stack exchange,提问作者Stoyan Milev
相关产品推荐
相关产品推荐

