You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.30 23:18:01