TF Slim:基于自定义数据集微调MobileNet V2遇形状匹配错误
解决Mobilenet_v2_1.4_224微调时的形状不匹配错误
你遇到的InvalidArgumentError是因为当前构建的模型和预训练检查点中的对应层参数形状不匹配,导致TensorFlow无法完成参数赋值。从错误信息里的lhs形状= [1,1,24,144] rhs形状= [1,1,32,192]来看,大概率是模型中某个卷积层的通道数和预训练检查点不一致,或是最后分类层的参数形状因自定义数据集类别数变化而不匹配。
下面是具体的排查和解决步骤:
1. 确认模型核心参数和预训练检查点完全一致
首先要确保你构建的Mobilenet_v2_1.4_224模型的关键参数和预训练检查点对齐:
- 宽度乘数(
depth_multiplier)必须设为1.4,这是Mobilenet_v2_1.4_224的核心参数,预训练检查点的所有层通道数都是基于这个乘数计算的,要是不小心用了1.0或0.75这类其他值,必然会出现通道数不匹配的问题。 - 输入图像尺寸必须是
224x224,和预训练模型的输入规格一致,避免特征图空间维度变化引发的形状冲突。 - 检查输入通道数:如果你的数据集不是RGB三通道(比如灰度图),模型输入层的通道数会和预训练的3通道不匹配,这时候要么把数据集转成三通道,要么单独处理输入层的初始化。
2. 排除不匹配的层,只加载可复用的预训练参数
因为你用的是自定义数据集做分类,最后一层分类层(Logits层)的参数形状肯定和预训练的1000类不一样,这时候需要在加载检查点时跳过这些层:
在你的训练代码中,修改加载检查点的部分,用slim.get_variables_to_restore指定要排除的层:
import tensorflow as tf from tensorflow.contrib import slim from nets import mobilenet_v2 # 构建自定义模型 num_classes = 你的数据集类别数 # 替换成你自己的类别数 images, labels = ... # 你的数据输入逻辑 logits, endpoints = mobilenet_v2.mobilenet_v2( images, num_classes=num_classes, is_training=True, depth_multiplier=1.4) # 务必确保这里的depth_multiplier是1.4 # 指定要加载的变量,排除最后分类相关的层 variables_to_restore = slim.get_variables_to_restore(exclude=[ 'MobilenetV2/Logits', 'MobilenetV2/Predictions' ]) # 从预训练检查点加载参数 checkpoint_path = 'path/to/your/mobilenet_v2_1.4_224.ckpt' init_fn = slim.assign_from_checkpoint_fn( checkpoint_path, variables_to_restore, ignore_missing_vars=False) # 要是你修改了中间层,可以设为True忽略缺失变量 # 在训练会话中初始化参数 with tf.Session() as sess: init_fn(sess) # 开始训练逻辑...
3. 检查是否修改了模型的中间结构
如果你对Mobilenet_v2的中间倒残差模块做了修改(比如调整通道数、增删模块),对应的预训练参数肯定无法匹配。这种情况下:
- 要么恢复模型的原始结构,只修改最后分类层;
- 要么在
exclude_scopes中添加所有修改过的层,让这些层随机初始化,只加载未修改的特征提取层参数。
4. 验证检查点和模型的变量名是否匹配
有时候变量名的细微差异也会导致加载失败,你可以用以下代码查看预训练检查点中的变量名和形状:
from tensorflow.python.tools.inspect_checkpoint import print_tensors_in_checkpoint_file print_tensors_in_checkpoint_file( file_name='path/to/mobilenet_v2_1.4_224.ckpt', tensor_name='', all_tensors=True, all_tensor_names=True)
然后对比当前模型的变量名,确认哪些变量形状不匹配,再针对性地排除它们。
内容的提问来源于stack exchange,提问作者Ravi
相关产品推荐
相关产品推荐

