TensorFlow 2.0 Keras训练CNN+GRU+CTC模型遇递归深度超限错误
你在TF2.x环境下迁移TF1.x的唇读模型时碰到的这个递归错误,我之前在迁移带CTC的序列模型时也遇到过——大概率是TF2.x和TF1.x的API差异,或是模型/生成器的实现兼容性问题,下面给你梳理几个核心排查点和解决方案:
1. 立刻替换fit_generator为fit
TF2.x已经完全弃用了fit_generator,虽然表面上还能调用,但内部实现逻辑和TF1.x差异极大,很容易触发递归或线程问题。直接改成fit即可,它现在原生支持生成器和tf.data.Dataset:
train_history = network_model.model.fit( lip_gen_train.get_batch(), steps_per_epoch=lip_gen_train.video_dataset_steps, epochs=1000, validation_data=lip_gen_test.get_batch(), validation_steps=lip_gen_test.video_dataset_steps, callbacks=[early_stop, checkpoint] # 你之前没把回调函数传进去,这也是个小问题! )
2. 修复CTC损失的实现逻辑
TF2.x中ctc_batch_cost的调用方式和TF1.x略有不同,你当前用的lambda y_true, y_pred: y_pred是个“假损失”,这种写法在TF1.x的静态图模式下能凑合用,但在TF2.x的动态图模式下很容易触发递归计算。建议换成标准的CTC损失函数:
def ctc_loss(y_true, y_pred): # 适配CTC输入要求:计算输入序列长度和标签长度 batch_len = 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) # 扩展长度张量适配batch维度 input_length = input_length * tf.ones(shape=(batch_len, 1), dtype=tf.int64) label_length = label_length * tf.ones(shape=(batch_len, 1), dtype=tf.int64) return tf.keras.backend.ctc_batch_cost(y_true, y_pred, input_length, label_length)
然后编译模型时替换成这个损失:
network_model.model.compile(loss={'ctc': ctc_loss}, optimizer=adam)
(另外CTC模型通常不需要accuracy指标,因为序列预测的准确率计算逻辑和普通分类不同,强行加可能也会引发额外问题)
3. 排查自定义生成器VideoGenerator的递归问题
TF2.x对生成器的线程安全要求更高,如果你在get_batch()或者build_data_from_frames()方法里有递归调用自身的逻辑,或者用了TF1.x的旧版队列/线程机制,很容易在TF2.x的动态图环境下触发无限递归。建议:
- 检查生成器内部是否有
while True之外的递归调用 - 把生成器改成返回
tf.data.Dataset的形式,TF2.x对Dataset的支持更稳定:
# 示例:把生成器转成Dataset def dataset_from_generator(generator, steps): return tf.data.Dataset.from_generator( lambda: generator, output_types=(your_input_type, your_label_type), output_shapes=(your_input_shape, your_label_shape) ).repeat().take(steps)
4. 移除TF1.x兼容代码,改用TF2.x原生GPU配置
你当前用的tf.compat.v1.ConfigProto和InteractiveSession在TF2.x里会干扰动态图的执行逻辑,换成TF2.x原生的GPU配置:
gpus = tf.config.list_physical_devices('GPU') if gpus: try: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) except RuntimeError as e: print(e)
5. 定位递归源头(终极调试手段)
如果上面的方法都没解决,用以下代码打印递归调用栈,找到触发无限递归的函数:
import traceback import sys # 设置一个合理的递归上限(别设太极端,否则会直接崩溃) sys.setrecursionlimit(10000) def trace_recursion(frame, event, arg): if event == 'call': current_stack = traceback.extract_stack() # 当递归深度接近上限时打印栈并退出 if len(current_stack) > sys.getrecursionlimit() - 20: print("\n=== 递归调用栈(接近上限) ===") traceback.print_stack() sys.exit(1) return trace_recursion # 开启递归跟踪 sys.settrace(trace_recursion)
把这段代码放在你的脚本开头,运行后就能看到到底是哪个函数在无限调用自身,针对性修复即可。
内容的提问来源于stack exchange,提问作者Souheil Fenghour

