如何在Android TF Lite中解码Keras手写识别模型的输出为字符串?
解码手写识别模型输出为文本字符串
输出维度含义
你的模型输出[1,32,81]对应:
1:批量大小(仅输入单张图片)32:序列长度(模型对输入图像的32个时间步做预测)81:字符类别数(包含空白符+训练时用到的所有字符,对应原Keras教程的字符集总数量)
解码步骤
要得到目标文本,需按照CTC解码规则处理输出,具体步骤如下:
1. 复现训练时的字符映射表
原Keras训练流程中会定义一个包含所有可识别字符的列表,还要额外加入一个空白符(通常对应索引0)。你需要在Android中创建完全一致的映射数组,示例如下:
// 需根据实际训练时的字符集调整,总长度必须为81 val charMap = arrayOf( "", // 索引0:CTC解码专用空白符 "0", "1", "2", "3", "4", "5", "6", "7", "8", "9", "a", "b", "c", "d", "e", "f", "g", ... // 补充所有训练时的字母、符号 )
2. 提取每个时间步的预测字符索引
遍历输出数组的32个时间步,对每个时间步的81个概率值取最大值对应的索引:
val outputArray = output[0] // 去掉batch维度,得到[32,81]的二维数组 val predictedIndices = mutableListOf<Int>() for (timeStep in outputArray) { var maxIndex = 0 var maxProb = timeStep[0] // 遍历当前时间步的所有类别概率,找到最大值索引 for (i in 1 until timeStep.size) { if (timeStep[i] > maxProb) { maxProb = timeStep[i] maxIndex = i } } predictedIndices.add(maxIndex) }
3. CTC解码:去重+移除空白符
按照CTC解码规则,去掉连续重复的字符和空白符,拼接成最终文本:
val stringBuilder = StringBuilder() var previousIndex = -1 for (index in predictedIndices) { // 跳过空白符,且跳过和前一个时间步重复的字符 if (index != 0 && index != previousIndex) { stringBuilder.append(charMap[index]) } previousIndex = index } val predictedText = stringBuilder.toString()
额外注意事项
- 输入预处理一致性:确保Bitmap转ByteBuffer的流程和训练时完全匹配:
- 原教程输入为32x128的灰度图,需先将彩色Bitmap转为灰度图,再缩放到对应尺寸
- 像素值需归一化到
[0,1](原教程用image = image / 255.0),转ByteBuffer时要将每个像素值除以255后转为float类型
- 字符映射表准确性:必须和训练模型时的字符集完全一致,否则解码结果会出现错乱
内容的提问来源于stack exchange,提问作者Mehdi Karbalai
相关产品推荐
相关产品推荐

