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

TensorFlow肿瘤分割训练数据耗尽中断问题求助

3D MRI肿瘤分割训练中断问题解决

问题背景

正在进行肿瘤分割任务,使用NIfTI格式的3D MRI图像,数据集共611张图像,尺寸为(240, 240, 160)。因3D图像内存占用过高无法全量加载至RAM,实现了数据生成器提取图像块,计算得训练集图像块数为61000、验证集为11375。

数据管道代码如下:

def load_nifti_image(filepath, patch_size=(48, 48, 32), step_size=(48, 48, 32)):
    nifti = nib.load(filepath)
    volume = nifti.get_fdata()

    # Create patches from the volume
    patches = patchify(volume, patch_size, step=step_size)

    # Reshape patches multiplying (5, 5, 5) and add channel dimension (1 for grayscale)
    patches = patches.reshape(-1, *patches.shape[-3:])
    patches = np.expand_dims(patches, axis=-1)

    return patches

# -----------------------TRAIN-----------------------
nifti_files = [os.path.join("/content/drive/MyDrive/Interpolated/train/images", f) for f in os.listdir("/content/drive/MyDrive/Interpolated/train/images") if f.endswith('.nii.gz')]
mask_files = [os.path.join("/content/drive/MyDrive/Interpolated/train/masks", f) for f in os.listdir("/content/drive/MyDrive/Interpolated/train/masks") if f.endswith('.nii.gz')]

 # -----------------------VALIDATION-----------------------
nifti_files_val = [os.path.join("/content/drive/MyDrive/Interpolated/validation/images", f) for f in os.listdir("/content/drive/MyDrive/Interpolated/validation/images") if f.endswith('.nii.gz')]
mask_files_val = [os.path.join("/content/drive/MyDrive/Interpolated/validation/masks", f) for f in os.listdir("/content/drive/MyDrive/Interpolated/validation/masks") if f.endswith('.nii.gz')]

def calculate_patches(filepath, patch_size=(48, 48, 32), step_size=(48, 48, 32)):
    nifti = nib.load(filepath)
    volume = nifti.get_fdata()

    # Calculate the number of patches
    patches_shape = [((i - p) // s) + 1 for i, p, s in zip(volume.shape, patch_size, step_size)]
    num_patches = np.prod(patches_shape)

    return num_patches

num_train_patches = sum(calculate_patches(f) for i, f in enumerate(nifti_files) if print(f"Processing file {i}...") is None)
num_val_patches = sum(calculate_patches(f) for i, f in enumerate(nifti_files_val) if print(f"Processing file {i}...") is None)

def data_generator(image_files, mask_files):
    for img_file, mask_file in zip(image_files, mask_files):
        image_patches = load_nifti_image(img_file)
        mask_patches = load_nifti_image(mask_file)

        for img_patch, mask_patch in zip(image_patches, mask_patches):
            yield img_patch, mask_patch

train_generator = data_generator(nifti_files, mask_files)
val_generator = data_generator(nifti_files_val, mask_files_val)

output_signature = (
    tf.TensorSpec(shape=(48, 48, 32, 1), dtype=tf.float64),
    tf.TensorSpec(shape=(48, 48, 32, 1), dtype=tf.float64)
)

dataset = tf.data.Dataset.from_generator(lambda: train_generator, output_signature=output_signature).repeat()
dataset_val = tf.data.Dataset.from_generator(lambda: val_generator, output_signature=output_signature).repeat()

dataset = dataset.batch(32)
dataset_val = dataset_val.batch(32)

test_model.fit(dataset, validation_data=dataset_val, epochs=100, steps_per_epoch=num_train_patches//32, validation_steps=num_val_patches//32)

训练至第2个epoch时出现以下警告并中断:

Epoch 1/100
1906/1906 [==============================] - 1963s 1s/step - loss: 0.6447 - dice_coefficient: 0.3553 - val_loss: 0.9113 - val_dice_coefficient: 0.0887
Epoch 2/100
   1/1906 [..............................] - ETA: 10:17 - loss: 1.0000 - dice_coefficient: 1.7961e-05

WARNING:tensorflow:Your input ran out of data; interrupting training. Make sure that your dataset or generator can generate at least `steps_per_epoch * epochs` batches (in this case, 190600 batches). You may need to use the repeat() function when building your dataset.
WARNING:tensorflow:Your input ran out of data; interrupting training. Make sure that your dataset or generator can generate at least `steps_per_epoch * epochs` batches (in this case, 355 batches). You may need to use the repeat() function when building your dataset.

1906/1906 [==============================] - 0s 31us/step - loss: 1.0000 - dice_coefficient: 1.7961e-05

<keras.src.callbacks.History at 0x7a793461a590>

问题原因

你定义的train_generator和val_generator是一次性迭代器,第一次epoch会耗尽所有数据,后续调用时已经没有数据可以生成。虽然在tf.data.Dataset上调用了.repeat(),但from_generator传入的lambda返回的是同一个已经耗尽的生成器实例,导致repeat无法重新生成数据。

解决方案

方案1:修改数据集创建逻辑,每次repeat时重新生成数据

把生成器的创建逻辑放到lambda内部,这样每次repeat的时候都会重新初始化一个新的生成器:

# 移除提前创建的train_generator和val_generator
# train_generator = data_generator(nifti_files, mask_files)
# val_generator = data_generator(nifti_files_val, mask_files_val)

output_signature = (
    tf.TensorSpec(shape=(48, 48, 32, 1), dtype=tf.float64),
    tf.TensorSpec(shape=(48, 48, 32, 1), dtype=tf.float64)
)

# 直接在lambda里调用data_generator创建新实例
dataset = tf.data.Dataset.from_generator(
    lambda: data_generator(nifti_files, mask_files),
    output_signature=output_signature
).repeat()

dataset_val = tf.data.Dataset.from_generator(
    lambda: data_generator(nifti_files_val, mask_files_val),
    output_signature=output_signature
).repeat()

dataset = dataset.batch(32)
dataset_val = dataset_val.batch(32)

test_model.fit(dataset, validation_data=dataset_val, epochs=100, steps_per_epoch=num_train_patches//32, validation_steps=num_val_patches//32)

方案2:让自定义生成器自身支持循环

修改data_generator函数,使其可以无限循环生成数据,这样无需依赖tf.data的repeat:

def data_generator(image_files, mask_files):
    while True:  # 无限循环
        # 每次循环打乱文件顺序,增加随机性
        paired_files = list(zip(image_files, mask_files))
        np.random.shuffle(paired_files)
        
        for img_file, mask_file in paired_files:
            image_patches = load_nifti_image(img_file)
            mask_patches = load_nifti_image(mask_file)
            
            # 打乱当前图像的patch顺序
            paired_patches = list(zip(image_patches, mask_patches))
            np.random.shuffle(paired_patches)
            
            for img_patch, mask_patch in paired_patches:
                yield img_patch, mask_patch

# 后续创建数据集时可以去掉repeat()
dataset = tf.data.Dataset.from_generator(
    lambda: data_generator(nifti_files, mask_files),
    output_signature=output_signature
).batch(32)

dataset_val = tf.data.Dataset.from_generator(
    lambda: data_generator(nifti_files_val, mask_files_val),
    output_signature=output_signature
).batch(32)

额外优化建议

  • 计算num_train_patches和num_val_patches时,无需重复加载NIfTI文件,可以直接从文件的元数据获取尺寸,避免IO开销:
def calculate_patches(filepath, patch_size=(48, 48, 32), step_size=(48, 48, 32)):
    nifti = nib.load(filepath)
    # 直接从header获取尺寸,不用加载整个数据
    volume_shape = nifti.header.get_data_shape()
    patches_shape = [((i - p) // s) + 1 for i, p, s in zip(volume_shape, patch_size, step_size)]
    num_patches = np.prod(patches_shape)
    return num_patches
  • 训练时对文件和patch进行打乱,能提升模型泛化能力,方案2中已经加入了这部分逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 14:06:02