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

迁移学习生成数据时遭遇TensorFlow Graph错误求助

多线程训练下TensorFlow图错误的排查与解决

你遇到的这个错误,核心原因是Keras多线程训练时,预训练模型的计算图无法在子线程中正确复用。从错误栈能看到,问题出在self._extractor.predict(frame_arr)这一步——当你用workers=4开启多线程后,每个工作线程会尝试创建自己的TensorFlow计算图,但你的base_model(InceptionV3)是在主线程里初始化的,子线程没法直接复用这个图,导致图初始化失败。

具体原因拆解

  • 你在主线程中初始化了预训练的Inception模型,把它传给了Extractor类;
  • fit_generator开启多线程后,Keras会启动多个子线程并行调用生成器;
  • 子线程里调用predict时,会尝试创建新的预测函数,但子线程没有继承主线程的图上下文,直接触发了图相关的错误。

解决方法

方法1:强制子线程复用主线程的图

这是最直接的修复方式,只需要修改Extractor类的代码:

  1. 在__init__方法中保存主线程的TensorFlow图:
# 记得先导入tensorflow
import tensorflow as tf

def __init__(self, path_to_data, seq_len, base_model, image_shape=(299, 299, 3)):
    super().__init__(path_to_data, seq_len, input_shape=image_shape)
    self._extractor = base_model
    # 保存主线程的计算图
    self._main_graph = tf.get_default_graph()
  1. 在调用predict的时候,用这个图的上下文包裹:
# 替换原来的features.append那一行
with self._main_graph.as_default():
    features.append(self._extractor.predict(frame_arr))

这样所有子线程都会复用主线程创建好的图,不会再出现图冲突的问题。

方法2:临时禁用多线程(用于调试)

如果只是想先验证代码逻辑是否正确,可以把fit_generator的workers参数设为0,强制在主线程运行生成器:

my_model.fit_generator(generator=train_gen, epochs=10, steps_per_epoch=steps_per_epoch, verbose=1, workers=0)

这个方法简单,但会降低训练速度,适合排查问题时使用。

方法3:预提取所有特征(更高效的迁移学习方案)

其实在迁移学习中,更高效的做法是先离线提取所有视频的特征,再用这些特征训练新模型,而不是在训练时实时提取:

  1. 先写一个脚本,遍历所有视频,用预训练模型提取每个帧的特征,保存成numpy文件;
  2. 训练新模型时直接加载这些预提取的特征,不需要在生成器里调用predict,彻底避免多线程图问题,还能节省重复提取特征的时间。

示例预提取代码:

import os
import numpy as np
from keras.preprocessing import image
from keras.applications.inception_v3 import preprocess_input

def pre_extract_features(extractor, save_dir):
    os.makedirs(save_dir, exist_ok=True)
    data = extractor.clean_data(extractor.get_data(), extractor._input_shape[0])
    class_one_hot = extractor.one_hot_encode()
    
    all_features = []
    all_labels = []
    
    for row in data:
        print(f"Processing {row[0].split('\\')[-1]}")
        sequence = extractor.get_frames(row[0])
        if len(sequence) > extractor._input_shape[0]:
            sequence = extractor.rescale_frame_list(sequence, extractor._input_shape[0])
        
        frame_features = []
        for frame in sequence:
            frame_arr = image.img_to_array(frame)
            frame_arr = np.expand_dims(frame_arr, axis=0)
            frame_arr = preprocess_input(frame_arr)
            # 用主线程的图提取特征
            with extractor._main_graph.as_default():
                feat = extractor._extractor.predict(frame_arr)
            frame_features.append(feat)
        
        all_features.append(frame_features)
        all_labels.append(class_one_hot[row[1]])
    
    # 保存特征和标签
    np.save(os.path.join(save_dir, "video_features.npy"), np.array(all_features))
    np.save(os.path.join(save_dir, "video_labels.npy"), np.array(all_labels))

之后训练时直接加载这些numpy文件,用普通的fit方法或者自定义生成器即可。

额外优化建议

  • 你的生成器用random.SystemRandom()随机选样本,容易导致样本重复或遗漏,建议用Keras的Sequence类实现线程安全的生成器,它天生支持多线程,比普通生成器更可靠;
  • scipy.misc.imresize已经被废弃了,建议换成cv2.resize或者skimage.transform.resize,避免后续版本的兼容性问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 08:11:35