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

TensorFlow 2.0 Keras训练CNN+GRU+CTC模型遇递归深度超限错误

解决TensorFlow 2.x中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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 06:57:38