Keras实现OCR时ctc_decode报错:Shape需为rank1却为rank2
解决Keras中CTCDecoder的Shape维度错误问题
你遇到的这个错误核心原因很明确:ctc_decode要求输入的序列长度张量是一维(rank 1),但你的input_x_widths是二维形状((batch_size, 1)),这就导致了CTCGreedyDecoder的维度不匹配问题。
解决起来很简单,只需要在传入CTC相关函数前,把长度张量的多余维度压缩掉就行。我们可以用keras.backend.squeeze(因为是TensorFlow后端,也可以用tf.squeeze)来实现这一点,下面是具体的修改步骤:
1. 修正CTC损失函数中的长度张量维度
首先,你的ctc_lambda_func里同样需要处理长度张量的维度,不然损失计算也会有潜在问题:
def ctc_lambda_func(args): y_pred, y_true, input_x_width, input_y_width = args # 压缩长度张量的最后一个维度,从(batch_size,1)变为(batch_size,) input_x_width = keras.backend.squeeze(input_x_width, axis=-1) input_y_width = keras.backend.squeeze(input_y_width, axis=-1) return keras.backend.ctc_batch_cost(y_true, y_pred, input_x_width, input_y_width)
2. 修正ctc_decode中的长度张量维度
然后修改你代码中解码器的定义部分,对input_x_widths做维度压缩:
# ctc decoder - 修改这里的长度张量处理 top_k_decoded, _ = keras.backend.ctc_decode( y_pred, keras.backend.squeeze(input_x_widths, axis=-1) # 去掉多余的维度 ) self.decoder = keras.backend.function([input_x, input_x_widths], [top_k_decoded[0]])
为什么这样改?
Keras中Input([1])定义的输入会生成形状为(batch_size, 1)的二维张量,但CTC相关的函数(ctc_batch_cost和ctc_decode)要求长度张量是一维的——每个元素对应一个样本的序列长度,形状为(batch_size,)。通过squeeze操作去掉最后一个维度后,就完全符合函数的输入要求了。
内容的提问来源于stack exchange,提问作者gdforj
相关产品推荐
相关产品推荐

