Keras验证码OCR模型训练时验证集迭代报InvalidArgumentError张量shape不匹配错误
问题根源
报错的核心原因是你的验证码标签长度不统一,tf.data.Dataset.batch() 操作要求同一个批次内的所有张量形状完全一致才能堆叠成批次张量。batch_size=1时不需要堆叠多份标签,所以运行正常;batch_size>1时遇到不同长度的标签就会触发形状不匹配错误,你收到的报错信息里[tensor]: [4], [batch]: [5]就说明同一个批次里同时存在长度为4和长度为5的标签。
排查步骤
首先全量检查数据集的标签长度分布,确认是否存在长度不一致的样本:
from collections import Counter # 统计训练集标签长度分布 train_label_lens = Counter([len(label) for label in y_train]) # 统计验证集标签长度分布 valid_label_lens = Counter([len(label) for label in y_valid]) print("训练集标签长度分布:", train_label_lens) print("验证集标签长度分布:", valid_label_lens)
解决方案
根据你的业务需求选择以下任意一种方案即可:
方案1:可变长度填充(推荐,适配CTC损失训练)
验证码OCR通常使用CTC损失,本身支持可变长度标签,你只需要把普通的batch操作替换为padded_batch,对不足长度的标签填充特殊标记即可:
修改你的数据集构建代码:
# 训练集修改 train_dataset = ( train_dataset.map( encode_single_sample, num_parallel_calls=tf.data.AUTOTUNE ) .padded_batch( batch_size=batch_size, padded_shapes={ "image": (img_width, img_height, 3), # 图片已经resize固定形状,无需填充 "label": (None,) # 标签按批次内最大长度自动填充 }, padding_values={ "image": 0.0, "label": char_to_num.vocab_size() # 用词汇表外的ID作为填充标记,对应CTC空白标记 } ) .prefetch(buffer_size=tf.data.AUTOTUNE).repeat(10) ) # 验证集同理修改 validation_dataset = ( validation_dataset.map( encode_single_sample, num_parallel_calls=tf.data.AUTOTUNE ) .padded_batch( batch_size=batch_size, padded_shapes={ "image": (img_width, img_height, 3), "label": (None,) }, padding_values={ "image": 0.0, "label": char_to_num.vocab_size() } ) .prefetch(buffer_size=tf.data.AUTOTUNE) )
如果使用该方案,打印标签时过滤掉填充ID再转字符即可正常展示。
方案2:过滤异常长度标签
如果你的业务要求验证码标签长度固定,直接过滤掉不符合长度要求的样本即可:
在数据集map操作后添加过滤逻辑,再调用普通batch:
# 假设要求标签固定长度为5,可根据实际需求修改 def filter_valid_label(sample): return tf.shape(sample["label"])[0] == 5 train_dataset = train_dataset.filter(filter_valid_label) validation_dataset = validation_dataset.filter(filter_valid_label)
内容的提问来源于stack exchange,提问作者Vinit Sutar
相关产品推荐
相关产品推荐

