使用ImageDataGenerator的class_mode='sparse'触发InvalidArgumentError的排查
问题原因与解决方案:ImageDataGenerator sparse模式下形状不兼容错误
问题原因
当设置class_mode='sparse'时,ImageDataGenerator输出的标签是一维整数张量(形状为(batch_size,)),而你在model.compile中使用的'precision'和'recall'是Keras的默认指标,它们默认期望输入的标签是独热编码的二维张量(形状为(batch_size, num_classes)),仅和class_mode='categorical'的输出格式匹配。
SparseCategoricalCrossentropy损失可以兼容稀疏标签,但默认的精度、召回率指标无法自动适配这种格式,因此计算指标时会触发形状不兼容的错误,也就是你看到的[1,256] vs. [1,64](本质是标签形状与指标期望的输入形状不匹配)。
解决方案
直接使用Keras提供的稀疏标签专用指标替换默认的字符串指标即可,具体修改model.compile中的metrics参数:
model.compile( optimizer = tf.keras.optimizers.Adam(learning_rate= 1e-4), loss = tf.keras.losses.SparseCategoricalCrossentropy(), metrics = [ 'accuracy', tf.keras.metrics.SparsePrecision(name='precision'), tf.keras.metrics.SparseRecall(name='recall') ] )
原理说明
SparsePrecision和SparseRecall是专门针对一维整数稀疏标签设计的指标,它们的输入格式与class_mode='sparse'生成的标签完全匹配,同时也能和SparseCategoricalCrossentropy损失协同工作,不会再出现形状不兼容的问题。
内容的提问来源于stack exchange,提问作者Darshil Pungalia
相关产品推荐
相关产品推荐

