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

TensorFlow多类别U-Net训练早停触发Python解释器状态错误

医学图像分割U-Net早停触发时的TensorFlow迭代器错误问题解决

我在用TensorFlow和Keras训练多类别医学图像分割U-Net模型时,碰到一个诡异的问题:只有当EarlyStopping(早停)触发训练停止时,才会抛出特定错误;如果手动提前停止训练,或者减少epoch数让训练在早停条件满足前结束,就完全没问题。

错误信息如下:

W tensorflow/core/kernels/data/generator_dataset_op.cc:108] Error occurred when finalizing GeneratorDataset iterator: FAILED_PRECONDITION: Python interpreter state is not initialized. The process may be terminated. [[{{node PyFunc}}]]

复现代码

import os
import numpy as np
import tensorflow as tf
from tensorflow import keras
from keras.models import Model
from keras.layers import Conv2D, MaxPooling2D, Dropout, UpSampling2D, concatenate
from keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau
from keras.preprocessing.image import ImageDataGenerator
from keras.utils import to_categorical


os.environ["CUDA_VISIBLE_DEVICES"] = "1"

# define the image size and batch size
img_size = (256, 256)
batch_size = 16
num_layers = 3
org_dir = r"somepathway"
mask_dir = r"somepathway"

seed = 3
epochs = 50
model_name = "test182v2"
num_classes = 5


def preprocess_masks(mask):
    conditions = [mask == 40, mask == 80, mask == 120, mask == 160]
    choices = [1, 2, 3, 4]
    # Apply np.select to set desired values based on conditions, and default to 0
    mask = np.select(conditions, choices, default=0)
    return to_categorical(mask, num_classes=num_classes)


# Modified generator to handle both images and masks
def combine_generator(img_gen, mask_gen):
    while True:
        try:
            img = img_gen.next()
            mask = mask_gen.next()
            yield (img, preprocess_masks(mask))
        except Exception as e:
            print(f"Error in generator: {e}")
            break

# Data generator setup
data_gen_args = dict(rescale=1./255,
                     validation_split=0.2)

image_datagen = ImageDataGenerator(rotation_range=20,
                                   width_shift_range=0.2,
                                   **data_gen_args)
mask_datagen = ImageDataGenerator(**data_gen_args)

# Create generators for images and masks
image_generator = image_datagen.flow_from_directory(
    org_dir, class_mode=None, color_mode='rgb',
    target_size=img_size, batch_size=batch_size, subset='training', seed=3)

mask_generator = mask_datagen.flow_from_directory(
    mask_dir, class_mode=None, color_mode='grayscale',
    target_size=img_size, batch_size=batch_size, subset='training', seed=3)

image_generator_val = image_datagen.flow_from_directory(
    org_dir, class_mode=None, color_mode='rgb',
    target_size=img_size, batch_size=batch_size, subset='validation', seed=3)

mask_generator_val = mask_datagen.flow_from_directory(
    mask_dir, class_mode=None, color_mode='grayscale',
    target_size=img_size, batch_size=batch_size, subset='validation', seed=3)

train_generator = combine_generator(image_generator, mask_generator)
validation_generator = combine_generator(image_generator_val, mask_generator_val)


def unet(input_size=(img_size[0], img_size[1], 3), num_classes=num_classes):
    inputs = keras.Input(input_size)

    # Encoder
    conv1 = Conv2D(64, (3, 3), activation='relu', padding='same')(inputs)
    conv1 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv1)
    pool1 = MaxPooling2D(pool_size=(2, 2))(conv1)

    # Defining the second level of the U-Net model
    conv2 = Conv2D(128, (3, 3), activation='relu', padding='same')(pool1)
    conv2 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv2)
    pool2 = MaxPooling2D(pool_size=(2, 2))(conv2)

    # Defining the third level of the U-Net model
    conv3 = Conv2D(256, (3, 3), activation='relu', padding='same')(pool2)
    conv3 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv3)
    pool3 = MaxPooling2D(pool_size=(2, 2))(conv3)

    # Defining the fourth level of the U-Net model
    conv4 = Conv2D(512, (3, 3), activation='relu', padding='same')(pool3)
    conv4 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv4)
    drop4 = Dropout(0.5)(conv4)
    pool4 = MaxPooling2D(pool_size=(2, 2))(drop4)

    # Defining the fifth level of the U-Net model
    conv5 = Conv2D(1024, (3, 3), activation='relu', padding='same')(pool4)
    conv5 = Conv2D(1024, (3, 3), activation='relu', padding='same')(conv5)
    drop5 = Dropout(0.5)(conv5)

    # Defining the sixth level of the U-Net model
    up6 = Conv2D(512, (2, 2), activation='relu', padding='same')(UpSampling2D(size=(2, 2))(drop5))
    merge6 = concatenate([drop4, up6], axis=3)
    conv6 = Conv2D(512, (3, 3), activation='relu', padding='same')(merge6)
    conv6 = Conv2D(512, (3, 3), activation='relu', padding='same')(conv6)

    # Defining the seventh level of the U-Net model
    up7 = Conv2D(256, (2, 2), activation='relu', padding='same')(UpSampling2D(size=(2, 2))(conv6))
    merge7 = concatenate([conv3, up7], axis=3)
    conv7 = Conv2D(256, (3, 3), activation='relu', padding='same')(merge7)
    conv7 = Conv2D(256, (3, 3), activation='relu', padding='same')(conv7)

    # Defining the eighth level of the
    up8 = Conv2D(128, (2, 2), activation='relu', padding='same')(UpSampling2D(size=(2, 2))(conv7))
    merge8 = concatenate([conv2, up8], axis=3)
    conv8 = Conv2D(128, (3, 3), activation='relu', padding='same')(merge8)
    conv8 = Conv2D(128, (3, 3), activation='relu', padding='same')(conv8)

    # Defining the ninth level of the U-Net model
    up9 = Conv2D(64, (2, 2), activation='relu', padding='same')(UpSampling2D(size=(2, 2))(conv8))
    merge9 = concatenate([conv1, up9], axis=3)
    conv9 = Conv2D(64, (3, 3), activation='relu', padding='same')(merge9)
    conv9 = Conv2D(64, (3, 3), activation='relu', padding='same')(conv9)

    # Defining the output layer of the U-Net model
    outputs = Conv2D(num_classes, (1, 1), activation='softmax')(conv9)

    # Defining the U-Net model
    model = Model(inputs=[inputs], outputs=[outputs])
    model.compile(optimizer=keras.optimizers.Adam(lr=0.0001), loss='categorical_crossentropy',
                  metrics=['accuracy', tf.keras.metrics.AUC()])
    model.summary()
    return model


checkpoint = ModelCheckpoint(filepath=f'{model_name}_weights.h5',
                             monitor='val_accuracy',
                             save_best_only=True,
                             save_weights_only=True,
                             verbose=1)

early_stopping = EarlyStopping(monitor='val_loss',
                               patience=2,
                               verbose=1,
                               restore_best_weights=True,
                               )

reduce_lr = ReduceLROnPlateau(monitor='val_loss',
                              factor=0.1,
                              patience=1,
                              verbose=1)

# define the model
model = unet(input_size=(img_size[0], img_size[1], 3), num_classes=num_classes)

# define the callbacks for real-time model performance tracking and evaluation
callbacks = [
    keras.callbacks.ModelCheckpoint(f'{model_name}.h5', save_best_only=True),
    keras.callbacks.TensorBoard(log_dir='./logs'),
    checkpoint, early_stopping, reduce_lr
]

# train the model
history = model.fit(train_generator, batch_size=batch_size, epochs=epochs, verbose=1,
                    callbacks=callbacks, validation_data=validation_generator,
                    steps_per_epoch=(len(os.listdir(r"somepathway"))*0.8)//batch_size,
                    validation_steps=(len(os.listdir(r"somepathway"))*0.2)//batch_size)

错误原因

这个问题本质是TensorFlow的异步迭代器和Python生成器之间的线程/进程同步冲突:

  • 早停回调触发时,TensorFlow会快速终止训练流程,但后台负责数据生成的线程仍在运行,且持有Python解释器的状态引用。
  • 当TensorFlow尝试清理迭代器资源时,后台生成器线程可能已被强制终止,导致Python解释器状态被销毁,进而触发"Python interpreter state is not initialized"错误。
  • 手动停止或正常结束训练时,TensorFlow会按流程等待当前epoch的所有数据生成完成后再清理资源,不会出现异步销毁的冲突。

解决方法

1. 替换为TensorFlow原生Dataset API(推荐)

放弃Keras的ImageDataGenerator和自定义生成器,改用tf.data.Dataset构建数据管道,它能更好地与TensorFlow的线程管理兼容,从根源避免同步问题:

# 示例:用tf.data加载图像和掩码
def load_image_mask(img_path, mask_path):
    # 读取并预处理图像
    img = tf.io.read_file(img_path)
    img = tf.image.decode_jpeg(img, channels=3)
    img = tf.image.resize(img, img_size)
    img = tf.cast(img, tf.float32) / 255.0
    
    # 读取并预处理掩码
    mask = tf.io.read_file(mask_path)
    mask = tf.image.decode_png(mask, channels=1)
    mask = tf.image.resize(mask, img_size, method=tf.image.ResizeMethod.NEAREST_NEIGHBOR)
    # 掩码值映射:40→1, 80→2, 120→3, 160→4,其余为0
    mask = tf.where(mask == 40, 1, 
                tf.where(mask == 80, 2, 
                    tf.where(mask == 120, 3, 
                        tf.where(mask == 160, 4, 0))))
    mask = tf.one_hot(tf.squeeze(mask), num_classes)
    return img, mask

# 构建训练数据集
train_img_dir = os.path.join(org_dir, "train")
train_mask_dir = os.path.join(mask_dir, "train")
train_img_paths = [os.path.join(train_img_dir, f) for f in os.listdir(train_img_dir)]
train_mask_paths = [os.path.join(train_mask_dir, f) for f in os.listdir(train_mask_dir)]

train_dataset = tf.data.Dataset.from_tensor_slices((train_img_paths, train_mask_paths))
train_dataset = train_dataset.map(load_image_mask, num_parallel_calls=tf.data.AUTOTUNE)
train_dataset = train_dataset.shuffle(buffer_size=100).batch(batch_size).prefetch(tf.data.AUTOTUNE)

# 构建验证数据集
val_img_dir = os.path.join(org_dir, "validation")
val_mask_dir = os.path.join(mask_dir, "validation")
val_img_paths = [os.path.join(val_img_dir, f) for f in os.listdir(val_img_dir)]
val_mask_paths = [os.path.join(val_mask_dir, f) for f in os.listdir(val_mask_dir)]

val_dataset = tf.data.Dataset.from_tensor_slices((val_img_paths, val_mask_paths))
val_dataset = val_dataset.map(load_image_mask, num_parallel_calls=tf.data.AUTOTUNE)
val_dataset = val_dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE)

# 训练时直接传入dataset
history = model.fit(train_dataset, epochs=epochs, verbose=1,
                    callbacks=callbacks, validation_data=val_dataset)

2. 给自定义生成器添加安全终止机制

如果不想替换生成器,可以修改生成器逻辑,添加全局终止标志,让生成器能优雅停止:

# 全局终止标志
stop_generator = False

def combine_generator(img_gen, mask_gen):
    global stop_generator
    while not stop_generator:
        try:
            img = img_gen.next()
            mask = mask_gen.next()
            yield (img, preprocess_masks(mask))
        except Exception as e:
            print(f"Error in generator: {e}")
            break

# 自定义早停回调,触发时设置终止标志
class SafeEarlyStopping(keras.callbacks.EarlyStopping):
    def on_train_end(self, logs=None):
        global stop_generator
        stop_generator = True
        super().on_train_end(logs)

# 替换原EarlyStopping为自定义的安全版本
early_stopping = SafeEarlyStopping(monitor='val_loss',
                                   patience=2,
                                   verbose=1,
                                   restore_best_weights=True)

3. 禁用异步迭代器(临时调试用)

在代码开头添加以下代码,强制TensorFlow同步执行数据生成,避免线程冲突,但会降低训练速度:

tf.data.experimental.enable_debug_mode()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 04:54:54