如何为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
相关产品推荐
相关产品推荐

