训练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任务特性,按以下步骤排查修复:
- 校验模型输出维度
CTC要求y_pred形状为(batch_size, max_sequence_length, num_classes),其中num_classes必须包含空白符(blank token)。检查模型最后一层输出维度,若字符类别为36种(26字母+10数字),则输出维度需设为37。 - 统一数据集维度
- 确保
train/test数据集的所有标签已做padding处理,y_true形状为(batch_size, max_label_length),长度完全统一。 - 确认模型输出的序列长度(如75)必须大于所有样本的标签长度,这是CTC计算的硬性要求。
- 确保
- 修正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 - 修复回调函数中硬编码的序列长度
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) - 兼容TensorFlow版本
针对tf.nn.ctc_loss废弃警告,可改用tf.compat.v1.nn.ctc_loss,或直接使用Keras内置的tf.keras.losses.CTCLoss(需适配新版本Keras)。
内容的提问来源于stack exchange,提问作者Mohsin Ali
相关产品推荐
相关产品推荐

