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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 19:38:10