使用自定义COCO数据集+TensorFlow数据生成器训练U-Net报错求助
问题描述
基于自定义COCO数据集训练U-Net模型时,使用自定义数据生成器函数触发了TensorFlow图执行错误。尝试调整输入/输出形状等方式排查多次无效,推测是逻辑问题,请求协助定位解决。
相关代码
数据生成器函数
def dataGeneratorCoco(images, classes, coco, folder, input_image_size=(224,224), batch_size=4, mode='train', mask_type='binary'): img_folder = '{}/{}-data/images'.format(folder,mode) dataset_size = len(images) catIds = coco.getCatIds(catNms=classes) c = 0 while(True): img = np.zeros((batch_size, input_image_size[0], input_image_size[1], 3)).astype('float') mask = np.zeros((batch_size, input_image_size[0], input_image_size[1], 1)).astype('float') for i in range(c, c+batch_size): #initially from 0 to batch_size, when c = 0 imageObj = images[i] ### Retrieve Image ### train_img = getImage(imageObj, img_folder, input_image_size) ### Create Mask ### if mask_type=="binary": train_mask = getBinaryMask(imageObj, coco, catIds, input_image_size) elif mask_type=="normal": train_mask = getNormalMask(imageObj, classes, coco, catIds, input_image_size) # Add to respective batch sized arrays # print(img.shape,train_img.shape) img[i-c] = train_img mask[i-c] = train_mask c+=batch_size if(c + batch_size >= dataset_size): c=0 random.shuffle(images) yield img, mask
数据增强生成器
def augmentationsGenerator(gen, augGeneratorArgs, seed=None): # Initialize the image data generator with args provided image_gen = ImageDataGenerator(**augGeneratorArgs) # Remove the brightness argument for the mask. Spatial arguments similar to image. augGeneratorArgs_mask = augGeneratorArgs.copy() _ = augGeneratorArgs_mask.pop('brightness_range', None) # Initialize the mask data generator with modified args mask_gen = ImageDataGenerator(**augGeneratorArgs_mask) np.random.seed(seed if seed is not None else np.random.choice(range(9999))) while(True): for img, mask in gen: seed = np.random.choice(range(9999)) # keep the seeds syncronized otherwise the augmentation of the images # will end up different from the augmentation of the masks g_x = image_gen.flow(255*img, batch_size = img.shape[0], seed = seed, shuffle=True) g_y = mask_gen.flow(mask, batch_size = mask.shape[0], seed = seed, shuffle=True) img_aug = next(g_x)/255.0 mask_aug = next(g_y) yield img_aug, mask_aug
报错信息
UnknownError: Graph execution error: 2 root error(s) found. (0) UNKNOWN: Exception: input type is not supported. Traceback (most recent call last): File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/ops/script_ops.py", line 270, in __call__ ret = func(*args) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/autograph/impl/api.py", line 642, in wrapper return func(*args, **kwargs) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/data/ops/dataset_ops.py", line 1030, in generator_py_func values = next(generator_state.get_iterator(iterator_id)) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/keras/engine/data_adapter.py", line 831, in wrapped_generator for data in generator_fn(): File "/tmp/ipykernel_284939/3839959985.py", line 15, in augmentationsGenerator for img, mask in gen: File "/tmp/ipykernel_284939/728327123.py", line 25, in dataGeneratorCoco train_mask = getNormalMask(imageObj, classes, coco, catIds, input_image_size) File "/tmp/ipykernel_284939/2567587607.py", line 29, in getNormalMask new_mask = cv2.resize(coco.annToMask(anns[a])*pixel_value, input_image_size) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/pycocotools/coco.py", line 442, in annToMask rle = self.annToRLE(ann) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/pycocotools/coco.py", line 427, in annToRLE rles = maskUtils.frPyObjects(segm, h, w) File "pycocotools/_mask.pyx", line 308, in pycocotools._mask.frPyObjects Exception: input type is not supported. [[{{node PyFunc}}]] [[IteratorGetNext]] [[Shape/_10]] (1) UNKNOWN: Exception: input type is not supported. Traceback (most recent call last): File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/ops/script_ops.py", line 270, in __call__ ret = func(*args) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/autograph/impl/api.py", line 642, in wrapper return func(*args, **kwargs) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/tensorflow/python/data/ops/dataset_ops.py", line 1030, in generator_py_func values = next(generator_state.get_iterator(iterator_id)) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/keras/engine/data_adapter.py", line 831, in wrapped_generator for data in generator_fn(): File "/tmp/ipykernel_284939/3839959985.py", line 15, in augmentationsGenerator for img, mask in gen: File "/tmp/ipykernel_284939/728327123.py", line 25, in dataGeneratorCoco train_mask = getNormalMask(imageObj, classes, coco, catIds, input_image_size) File "/tmp/ipykernel_284939/2567587607.py", line 29, in getNormalMask new_mask = cv2.resize(coco.annToMask(anns[a])*pixel_value, input_image_size) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/pycocotools/coco.py", line 442, in annToMask rle = self.annToRLE(ann) File "/home/navneeth/anaconda3/lib/python3.9/site-packages/pycocotools/coco.py", line 427, in annToRLE rles = maskUtils.frPyObjects(segm, h, w) File "pycocotools/_mask.pyx", line 308, in pycocotools._mask.frPyObjects Exception: input type is not supported. [[{{node PyFunc}}]] [[IteratorGetNext]] 0 successful operations. 0 derived errors ignored. [Op:__inference_train_function_3047]
解决方案
从报错堆栈看,问题出在getNormalMask函数调用coco.annToMask(anns[a])时,传入的标注格式不符合pycocotools要求,具体排查和修复方向如下:
- 检查标注数据格式:确认自定义COCO数据集中的
segmentation字段类型是否合法。pycocotools的frPyObjects仅接受三种输入:单个多边形(list of list of ints)、多个多边形(list of list of list of ints)、RLE格式字典(含counts和size键)。检查anns[a]的segmentation是否为None、字符串或其他非法类型。 - 修复标注遍历逻辑:在
getNormalMask中添加格式校验,跳过或处理无效标注:for a in range(len(anns)): ann = anns[a] segm = ann.get('segmentation') if not isinstance(segm, (list, dict)): print(f"Invalid segmentation for annotation {ann['id']}: {type(segm)}") continue # 后续mask生成逻辑 - 过滤无效图像:检查数据集中是否存在无标注或标注为空的图像,在数据生成器中提前过滤这类样本,避免传入无效标注给pycocotools。
- 兼容TensorFlow图模式:若报错仍在图执行阶段出现,可将pycocotools相关调用用
tf.py_function包装,确保Python代码在图模式下能正常执行,但优先解决标注格式问题。
内容的提问来源于stack exchange,提问作者Navneeth S
相关产品推荐
相关产品推荐

