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
相关产品推荐
相关产品推荐

