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

如何为Keras手写识别OCR模型添加decode_batch_predictions()方法?

适配Keras手写识别OCR模型的decode_batch_predictions集成方案

核心思路

把CTC解码逻辑封装成自定义Keras层,将原模型的输出连接到这个解码层,形成端到端模型。转换为TF Lite后,模型就能直接输出解码后的文本结果,无需在Android端额外处理。

具体实现代码

1. 定义CTC解码自定义层

该层基于TensorFlow可追踪操作实现,兼容TF Lite转换,同时完成从logits到文本的解码:

import tensorflow as tf
from tensorflow import keras

class CTCDecoderLayer(keras.layers.Layer):
    def __init__(self, char_to_num, **kwargs):
        super().__init__(**kwargs)
        self.char_to_num = char_to_num
        # 生成数字到字符的映射表
        self.num_to_char = tf.keras.layers.StringLookup(
            vocabulary=char_to_num.get_vocabulary(),
            invert=True,
            mask_token=None
        )

    def call(self, inputs):
        # 贪心解码:取每个时间步概率最大的索引
        indices = tf.argmax(inputs, axis=-1)
        # 转换为字符序列
        chars = self.num_to_char(indices)
        
        # 实现CTC解码核心逻辑:去空白符、合并连续重复字符
        def decode_single(chars_seq):
            # 过滤空白符
            filtered = tf.gather(chars_seq, tf.where(chars_seq != ''))
            filtered = tf.squeeze(filtered, axis=1)
            # 合并连续重复字符
            unique_idx = tf.where(tf.not_equal(filtered[1:], filtered[:-1]))
            unique_idx = tf.concat([[0], unique_idx[:, 0] + 1], axis=0)
            unique_chars = tf.gather(filtered, unique_idx)
            # 拼接成字符串
            return tf.strings.reduce_join(unique_chars)
        
        # 批量处理样本
        decoded_texts = tf.map_fn(decode_single, chars, dtype=tf.string)
        return decoded_texts

    def get_config(self):
        config = super().get_config()
        config.update({"char_to_num": self.char_to_num})
        return config

2. 集成到原模型

假设你已完成官方示例模型的训练,按以下步骤构建端到端模型:

# 假设`model`是训练好的原模型,`char_to_num`是官方示例中定义的字符映射层
decoder = CTCDecoderLayer(char_to_num)
# 构建输入到解码文本的端到端模型
end_to_end_model = keras.Model(
    inputs=model.input,
    outputs=decoder(model.output)
)

# 测试输出:直接得到解码后的文本
# test_images为预处理后的测试图片(格式与训练时一致)
predictions = end_to_end_model.predict(test_images)
print(predictions)  # 输出文本数组,如["hello", "world"]

3. 转换为TF Lite模型

# 初始化转换器
converter = tf.lite.TFLiteConverter.from_keras_model(end_to_end_model)
# 启用TF扩展操作,确保自定义层逻辑可转换
converter.target_spec.supported_ops = [
    tf.lite.OpsSet.TFLITE_BUILTINS,
    tf.lite.OpsSet.SELECT_TF_OPS
]
# 转换并保存模型
tflite_model = converter.convert()
with open("handwriting_ocr_decoded.tflite", "wb") as f:
    f.write(tflite_model)

关键说明

  • 采用贪心解码而非Beam Search,因为后者在TF Lite中支持有限,贪心解码足以覆盖多数手写识别场景,且转换过程更稳定。
  • 自定义层逻辑与官方示例的decode_batch_predictions完全对齐,保证解码结果一致。
  • 启用SELECT_TF_OPS是为了兼容tf.map_fn等操作,避免转换失败。

内容的提问来源于stack exchange,提问作者Mehdi Karbalai

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 19:55:22