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

手写文本识别模型CTC损失层报错:indices参数数据类型不符

问题描述

在构建包含2-D CNN、Bidirectional LSTM的手写文本识别模型时,自定义CTC损失层遇到以下错误:

TypeError: 调用CTC_Layer.call()时遇到异常。
Value passed to parameter 'indices' has DataType float32 not in list of allowed values: uint8, int8, int32, int64

自定义CTC_Layer代码:

class CTC_Layer(Layer):
    def __init__(self, name=None):
        super(CTC_Layer, self).__init__(name='ctc_loss')
        self.loss_fn = tensorflow.nn.ctc_loss
        
    def call(self, y_true, y_pred):
        batch_length = tf.cast(tf.shape(y_true)[0], "int64")
        input_length = tf.cast(tf.shape(y_pred)[1], "int64")
        label_length = tf.cast(tf.shape(y_true)[1], "int64")
        
        input_length = input_length * tf.ones(shape=(batch_length,), dtype="int64")
        label_length = label_length * tf.ones(shape=(batch_length,), dtype="int64")
        
        loss = self.loss_fn(y_true, y_pred, input_length, label_length)
        self.add_loss(loss)
        
        return y_pred

报错堆栈信息:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
Cell In[133], line 1
----> 1 model = build_model()
      2 model.summary()

Cell In[132], line 31, in build_model()
     24 x = layers.Bidirectional(
     25     layers.LSTM(64, return_sequences=True, dropout=0.25)
     26 )(x)
     28 x = layers.Dense(
     29     len((char_to_num.get_vocabulary()))+2, activation='softmax', name='dense2'
     30 )(x)
---> 31 outputs = CTC_Layer(name="ctc_loss")(labels, x)
     34 model = keras.models.Model(inputs=[input_image, labels], outputs=outputs)
     36 model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.001))

File /opt/conda/lib/python3.10/site-packages/keras/src/utils/traceback_utils.py:122, in filter_traceback.<locals>.error_handler(*args, **kwargs)
    119     filtered_tb = _process_traceback_frames(e.__traceback__)
    120     # To get the full stack trace, call:
    121     # `keras.config.disable_traceback_filtering()`
--> 122     raise e.with_traceback(filtered_tb) from None
    123 finally:
    124     del filtered_tb

Cell In[93], line 15, in CTC_Layer.call(self, y_true, y_pred)
     12 input_length = input_length * tf.ones(shape=(batch_length,), dtype="int64")
     13 label_length = label_length * tf.ones(shape=(batch_length,), dtype="int64")
---> 15 loss = self.loss_fn(y_true, y_pred, input_length, label_length)
     16 self.add_loss(loss)
     18 # At test time, just return the computed predictions.

TypeError: Exception encountered when calling CTC_Layer.call().

Value passed to parameter 'indices' has DataType float32 not in list of allowed values: uint8, int8, int32, int64

Arguments received by CTC_Layer.call():
  • args=('<KerasTensor shape=(None, None), dtype=float32, sparse=None, name=label>', '<KerasTensor shape=(None, 32, 79), dtype=float32, sparse=False, name=keras_tensor_317>')
  • kwargs=<class 'inspect._empty'>]
解决方案

错误核心是tf.nn.ctc_loss要求输入的标签(y_true)必须是整数类型(uint8/int8/int32/int64),但当前传入的标签是float32类型,需按以下方式修复:

1. 修改CTC层代码,强制转换标签类型

在call方法开头添加类型转换代码,将float32的标签转为整数类型:

class CTC_Layer(Layer):
    def __init__(self, name=None):
        super(CTC_Layer, self).__init__(name='ctc_loss')
        self.loss_fn = tensorflow.nn.ctc_loss
        
    def call(self, y_true, y_pred):
        # 关键修复:将浮点型标签转为int32
        y_true = tf.cast(y_true, tf.int32)
        
        batch_length = tf.cast(tf.shape(y_true)[0], tf.int64)
        input_length = tf.cast(tf.shape(y_pred)[1], tf.int64)
        label_length = tf.cast(tf.shape(y_true)[1], tf.int64)
        
        input_length = input_length * tf.ones(shape=(batch_length,), dtype=tf.int64)
        label_length = label_length * tf.ones(shape=(batch_length,), dtype=tf.int64)
        
        loss = self.loss_fn(y_true, y_pred, input_length, label_length)
        self.add_loss(loss)
        
        return y_pred

2. 检查数据预处理环节

确认数据生成或预处理时,标签没有被错误转换为浮点型:

  • 检查char_to_num层的输出类型,确保它输出整数张量
  • 避免对标签进行归一化等会改变数据类型的操作

内容的提问来源于stack exchange,提问作者Manish Kumar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 00:37:44