如何修改Keras验证码OCR模型使其支持可变长度样本输入输出
报错原因
你遇到的报错是因为tf.data.Dataset默认的batch方法要求批次内所有样本的张量形状完全一致,当你把部分文件名改成更长的字符串后,对应标签的编码长度从5位变成10位,和其他样本的长度不匹配,就触发了维度不匹配的错误。
可变长度适配修改方案
无需提前固定全局最大标签长度,通过动态填充+CTC损失自适应计算即可实现任意长度标签的支持,仅需修改3处代码即可:
1. 调整数据集批次生成逻辑
将原来的普通batch方法替换为padded_batch,对可变长度的标签进行动态填充:
def create_dataset(self, x, y, batch_size): dataset = tf.data.Dataset.from_tensor_slices((x, y)) return ( dataset.map(self.encode_sample, num_parallel_calls=tf.data.AUTOTUNE) .padded_batch( batch_size=batch_size, padded_shapes={ "image": (self.image_width, self.image_height, 1), # 图片尺寸固定无需填充 "label": (None,) # 标签按当前批次内最长长度动态填充 }, padding_values={ "image": tf.constant(0, dtype=tf.float32), "label": tf.constant(-1, dtype=tf.int64) # 用-1作为填充标记,后续损失计算会忽略 } ) .prefetch(buffer_size=tf.data.AUTOTUNE) )
2. 修改CTC损失层的标签长度计算逻辑
原来的标签长度取全局批次的第二维长度,现在改为统计每个样本的非填充值数量,得到实际标签长度:
def call(self, y_true, *args, **kwargs): y_pred = args[0] batch_length = tf.cast(tf.shape(y_true)[0], dtype='int64') input_length = tf.cast(tf.shape(y_pred)[1], dtype='int64') input_length = input_length * tf.ones(shape=(batch_length, 1), dtype='int64') # 替换原来的固定长度计算,统计非-1的元素作为实际标签长度 label_length = tf.math.count_nonzero(y_true != -1, axis=1, keepdims=True, dtype=tf.int64) loss = self.loss_fn(y_true, y_pred, input_length, label_length) self.add_loss(loss) return y_pred
3. 调整预测解码逻辑
去掉原来固定长度的截断限制,适配可变长度输出:
def decode_batch_predictions(self, pred): input_len = np.ones(pred.shape[0]) * pred.shape[1] results = tf.keras.backend.ctc_decode( pred, input_length=input_len, greedy=True )[0][0] # 去掉原来的[:, : self.max_label_length]截断逻辑 output_text = [] for result in results: # 过滤掉填充对应的空字符 result = tf.boolean_mask(result, result != -1) result = ( tf.strings.reduce_join(self.num_to_char(result)).numpy().decode('utf-8') ) output_text.append(result) return output_text
方案说明
该方案不存在提前固定最大长度的限制:
- 每个批次的标签会自动按当前批次内最长的标签长度填充,新增更长标签样本无需修改代码
- 唯一的限制来自模型本身的时间步长度:当前模型经过两次2倍池化后,时间步长度为图片宽度/4=50,只要标签长度不超过50都可以正常支持,如果需要更长的输出,加宽输入图片或者减少池化层数即可。
内容的提问来源于stack exchange,提问作者watch-this
相关产品推荐
相关产品推荐

