MaskRCNN测试Transforms时触发NotImplementedError问题求助
问题:MaskRCNN教程中Test the transforms步骤触发NotImplementedError
场景
学习MaskRCNN时,按教程步骤操作,前面环节均正常,执行"Test the transforms"代码块时触发错误。
执行代码
# Extract the labels for the sample labels = [shape['label'] for shape in annotation_df.loc[file_id]['shapes']] # Extract the polygon points for segmentation mask shape_points = [shape['points'] for shape in annotation_df.loc[file_id]['shapes']] # Format polygon points for PIL xy_coords = [[tuple(p) for p in points] for points in shape_points] # Generate mask images from polygons mask_imgs = [create_polygon_mask(sample_img.size, xy) for xy in xy_coords] # Convert mask images to tensors masks = torch.concat([Mask(transforms.PILToTensor()(mask_img), dtype=torch.bool) for mask_img in mask_imgs]) # Generate bounding box annotations from segmentation masks bboxes = BoundingBoxes(data=torchvision.ops.masks_to_boxes(masks), format='xyxy', canvas_size=sample_img.size[::-1]) # Get colors for dataset sample sample_colors = [int_colors[i] for i in [class_names.index(label) for label in labels]] # Prepare mask and bounding box targets targets = { 'masks': Mask(masks), 'boxes': bboxes, 'labels': torch.Tensor([class_names.index(label) for label in labels]) } # Crop the image cropped_img, targets = iou_crop(sample_img, targets) # Resize the image resized_img, targets = resize_max(cropped_img, targets) # Pad the image padded_img, targets = pad_square(resized_img, targets) # Ensure the padded image is the target size resize = transforms.Resize([train_sz] * 2, antialias=True) resized_padded_img, targets = resize(padded_img, targets) sanitized_img, targets = transforms.SanitizeBoundingBoxes()(resized_padded_img, targets) # Annotate the sample image with segmentation masks annotated_tensor = draw_segmentation_masks( image=transforms.PILToTensor()(sanitized_img), masks=targets['masks'], alpha=0.3, colors=sample_colors ) # Annotate the sample image with labels and bounding boxes annotated_tensor = draw_bboxes( image=annotated_tensor, boxes=targets['boxes'], labels=[class_names[int(label.item())] for label in targets['labels']], colors=sample_colors ) # # Display the annotated image display(tensor_to_pil(annotated_tensor)) pd.Series({ "Source Image:": sample_img.size, "Cropped Image:": cropped_img.size, "Resized Image:": resized_img.size, "Padded Image:": padded_img.size, "Resized Padded Image:": resized_padded_img.size, }).to_frame().style.hide(axis='columns')
报错信息
--------------------------------------------------------------------------- NotImplementedError Traceback (most recent call last) Cell In[33], line 28 25 cropped_img, targets = iou_crop(sample_img, targets) 27 # Resize the image ---> 28 resized_img, targets = resize_max(cropped_img, targets) 30 # Pad the image 31 padded_img, targets = pad_square(resized_img, targets) File ~\anaconda3\envs\test-env\lib\site-packages\torch\nn\modules\module.py:1739, in Module._wrapped_call_impl(self, *args, **kwargs) 1737 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc] 1738 else: -> 1739 return self._call_impl(*args, **kwargs) File ~\anaconda3\envs\test-env\lib\site-packages\torch\nn\modules\module.py:1750, in Module._call_impl(self, *args, **kwargs) 1745 # If we don't have any hooks, we want to skip the rest of the logic in 1746 # this function, and just call forward. 1747 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks 1748 or _global_backward_pre_hooks or _global_backward_hooks 1749 or _global_forward_hooks or _global_forward_pre_hooks): -> 1750 return forward_call(*args, **kwargs) 1752 result = None 1753 called_always_called_hooks = set() File ~\anaconda3\envs\test-env\lib\site-packages\torchvision\transforms\v2\_transform.py:68, in Transform.forward(self, *inputs) 63 needs_transform_list = self._needs_transform_list(flat_inputs) 64 params = self.make_params( 65 [inpt for (inpt, needs_transform) in zip(flat_inputs, needs_transform_list) if needs_transform] 66 ) ---> 68 flat_outputs = [ 69 self.transform(inpt, params) if needs_transform else inpt 70 for (inpt, needs_transform) in zip(flat_inputs, needs_transform_list) 71 ] 73 return tree_unflatten(flat_outputs, spec) File ~\anaconda3\envs\test-env\lib\site-packages\torchvision\transforms\v2\_transform.py:69, in <listcomp>(.0) 63 needs_transform_list = self._needs_transform_list(flat_inputs) 64 params = self.make_params( 65 [inpt for (inpt, needs_transform) in zip(flat_inputs, needs_transform_list) if needs_transform] 66 ) 68 flat_outputs = [ ---> 69 self.transform(inpt, params) if needs_transform else inpt 70 for (inpt, needs_transform) in zip(flat_inputs, needs_transform_list) 71 ] 73 return tree_unflatten(flat_outputs, spec) File ~\anaconda3\envs\test-env\lib\site-packages\torchvision\transforms\v2\_transform.py:55, in Transform.transform(self, inpt: Any, params: Dict[str, Any]) -> Any: 51 def transform(self, inpt: Any, params: Dict[str, Any]) -> Any: 52 """Method to override for custom transforms. 53 54 See :ref:`sphx_glr_auto_examples_transforms_plot_custom_transforms.py`""" ---> 55 raise NotImplementedError NotImplementedError:
问题定位
错误根源是torchvision.transforms.v2._transform.Transform类的transform方法未被实现。该方法是v2版本自定义Transform的必填重写方法,直接调用基类会抛出此错误。
解决思路
- 检查
resize_max的定义:如果resize_max是自定义的Transform类,必须确保它继承自torchvision.transforms.v2.Transform后,正确重写了make_params和transform方法,实现具体的缩放逻辑(比如计算长边缩放比例、调整图像尺寸,同时更新targets中的boxes和masks坐标)。 - 核对torchvision版本:若教程使用的是旧版torchvision(v1.x),而你安装的是v2.x,两者Transform API差异较大,需按v2的规范修改自定义Transform代码,或者降级torchvision到教程对应的版本。
- 临时替代方案:若自定义
resize_max实现有问题,可手动实现相同逻辑:计算图像的缩放比例,保证长边不超过目标尺寸,用torchvision内置的Resize转换图像,再手动调整targets中的boxes和masks(按缩放比例修改坐标/尺寸)。
内容的提问来源于stack exchange,提问作者Rhjg
相关产品推荐
相关产品推荐

