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

训练700轮时出现Graph Execution/Invalid Argument错误求助

CTC训练中Graph Execution/Invalid Argument错误解决

问题描述

在运行OCR相关的CTC训练代码时,已解决多个前置错误,但执行训练轮次时触发Graph Execution/Invalid Argument错误,即使仅测试10轮仍报错。

训练代码

def scheduler(epoch, lr):
    if epoch < 30:
        return lr
    else:
        return lr * tf.math.exp(-0.1)

# 采用CTC LOSS函数计算损失,因其在图像和视频序列任务中表现出色
def CTCLoss(y_true, y_pred):
    batch_len = tf.cast(tf.shape(y_true)[0], dtype="int64")
    input_length = tf.cast(tf.shape(y_pred)[1], dtype="int64")
    label_length = tf.cast(tf.shape(y_true)[1], dtype="int64")

    input_length = input_length * tf.ones(shape=(batch_len, 1), dtype="int64")
    label_length = label_length * tf.ones(shape=(batch_len, 1), dtype="int64")

    loss = tf.keras.backend.ctc_batch_cost(y_true, y_pred, input_length, label_length)
    return loss

# 通过类方法生成预测,包含数字与字符互转功能
class ProduceExample(tf.keras.callbacks.Callback):
    def __init__(self, dataset) -> None:
        self.dataset = dataset.as_numpy_iterator()

    def on_epoch_end(self, epoch, logs=None) -> None:
        data = self.dataset.next()
        yhat = self.model.predict(data[0])
        decoded = tf.keras.backend.ctc_decode(yhat, [75,75], greedy=False)[0][0].numpy()
        for x in range(len(yhat)):
            print('Original:', tf.strings.reduce_join(num_to_char(data[1][x])).numpy().decode('utf-8'))
            print('Prediction:', tf.strings.reduce_join(num_to_char(decoded[x])).numpy().decode('utf-8'))
            print('~'*100)

# 使用Adam优化器编译模型
model.compile(optimizer=Adam(learning_rate=0.0001), loss=CTCLoss)

checkpoint_callback = ModelCheckpoint(os.path.join('models','checkpoint'), monitor='loss', save_weights_only=True)

schedule_callback = LearningRateScheduler(scheduler)

example_callback = ProduceExample(test)

# 执行以下代码时触发错误
model.fit(train, validation_data=test, epochs=700, callbacks=[checkpoint_callback, schedule_callback, example_callback])

报错信息

Epoch 1/10
WARNING:tensorflow:From e:\Pyhton 3.12\Lib\site-packages\keras\src\legacy\backend.py:666: The name tf.nn.ctc_loss is deprecated. Please use tf.compat.v1.nn.ctc_loss instead.
<hr />
InvalidArgumentError                      Traceback (most recent call last)
Cell In[108], line 1
----> 1 model.fit(train,validation_data=test ,epochs=10 ,callbacks =[checkpoint_callback, schedule_callback, example_callback])
...
File "e:\Pyhton 3.12\Lib\site-packages\keras\src\backend\tensorflow\numpy.py", line 1618, in reshape

Only one input size may be -1, not both 0 and 1
[[{{node sequential_11_1/time_distributed_9_1/Reshape_72}}]] [Op:__inference_one_step_on_iterator_49320]

错误原因与解决办法

  • 核心根源:Reshape节点报错,说明模型输出或数据集的维度不匹配,导致reshape操作时出现两个不确定维度(0和-1)。结合CTC任务特性,按以下步骤排查修复:
  1. 校验模型输出维度
    CTC要求y_pred形状为(batch_size, max_sequence_length, num_classes),其中num_classes必须包含空白符(blank token)。检查模型最后一层输出维度,若字符类别为36种(26字母+10数字),则输出维度需设为37。
  2. 统一数据集维度
    • 确保train/test数据集的所有标签已做padding处理,y_true形状为(batch_size, max_label_length),长度完全统一。
    • 确认模型输出的序列长度(如75)必须大于所有样本的标签长度,这是CTC计算的硬性要求。
  3. 修正CTC Loss函数的维度处理
    替换原Loss函数,确保input_length和label_length维度正确:
    def CTCLoss(y_true, y_pred):
        # y_pred形状:(batch_size, seq_len, num_classes)
        batch_len = tf.shape(y_true)[0]
        input_length = tf.fill((batch_len,), tf.shape(y_pred)[1])
        label_length = tf.fill((batch_len,), tf.shape(y_true)[1])
        
        input_length = tf.cast(input_length, tf.int64)
        label_length = tf.cast(label_length, tf.int64)
        
        loss = tf.keras.backend.ctc_batch_cost(y_true, y_pred, input_length, label_length)
        return loss
    
  4. 修复回调函数中硬编码的序列长度
    ProduceExample回调里硬编码的[75,75]仅适用于batch size=2的场景,改为动态获取:
    def on_epoch_end(self, epoch, logs=None) -> None:
        data = self.dataset.next()
        yhat = self.model.predict(data[0])
        # 动态生成输入长度
        input_len = tf.fill((tf.shape(yhat)[0],), tf.shape(yhat)[1])
        decoded = tf.keras.backend.ctc_decode(yhat, input_len, greedy=False)[0][0].numpy()
        for x in range(len(yhat)):
            print('Original:', tf.strings.reduce_join(num_to_char(data[1][x])).numpy().decode('utf-8'))
            print('Prediction:', tf.strings.reduce_join(num_to_char(decoded[x])).numpy().decode('utf-8'))
            print('~'*100)
    
  5. 兼容TensorFlow版本
    针对tf.nn.ctc_loss废弃警告,可改用tf.compat.v1.nn.ctc_loss,或直接使用Keras内置的tf.keras.losses.CTCLoss(需适配新版本Keras)。

内容的提问来源于stack exchange,提问作者Mohsin Ali

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 13:55:05