基于批量数据集的CNN模型形状不兼容问题排查
解决CNN训练时的形状不兼容错误
嘿,这个错误我之前也碰到过,咱们一步步来理清楚:
先搞懂(None,1)是哪来的
你用tf.keras.preprocessing.image_dataset_from_directory加载数据集时,默认的标签模式是label_mode='int'——也就是说,数据集给你的标签是单个整数(比如0到4,对应5个类别),所以每个批次的标签形状是(批量大小, 1),也就是报错里的(None,1)(None代表批量大小不固定)。
而你的模型最后一层是Dense(5),输出的是5个类别的概率分布,形状是(批量大小,5),这俩形状对不上,自然就会抛出不兼容的错误。
两种快速解决的办法,选一种就行
办法1:让数据集输出独热编码标签
直接在加载数据集的时候,把label_mode改成'categorical',这样标签会变成5维的独热向量(比如类别0对应[1,0,0,0,0]),形状就变成(None,5),和模型输出完美匹配:
train_ds = tf.keras.preprocessing.image_dataset_from_directory( data_dir, seed=123, image_size=(img_height, img_width), label_mode='categorical' # 加这一行就搞定 )
这种情况下,你训练时用的损失函数如果是CategoricalCrossentropy(默认的多分类损失),就不用改了。
办法2:保留整数标签,修改损失函数
如果你不想改数据集的标签模式,那就要调整模型的损失函数:
- 首先,确保模型最后一层加上
activation='softmax'(虽然不是必须,但多分类任务加了更合理,能输出概率分布):
Dense(5, activation='softmax')
- 然后,训练时把损失函数换成
SparseCategoricalCrossentropy——这个损失函数就是专门用来处理整数标签和多分类输出的匹配问题的:
model.compile( optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(), metrics=['accuracy'] )
小提醒
如果你之前已经编译过模型,修改完损失函数或者模型层之后,一定要重新编译再运行model.fit()哦!
内容的提问来源于stack exchange,提问作者Davide Maran
相关产品推荐
相关产品推荐

