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

PyTorch数据集新增图像显示异常问题排查求助

问题描述

在数据集内添加新的彩色图像后,调用test_dataset.show_data(2, show_annotation=False)时显示异常,而旧图像调用test_dataset.show_data(12, show_annotation=False)显示正常。已注释__getitem__方法中的resize代码,也尝试跳过image_ins = image_ins.convert('RGB')语句,但问题仍未解决。

附数据集类的三个方法代码:

# Getter
def __getitem__(self, index):
    # load image and return it as tensor
    fn = os.path.join(self.images_path, self.im.file_name[index])
    orig_width  = self.width
    orig_height = self.height 
    if os.path.exists(fn):
        image_ins = Image.open(fn)
        """
        orig_width  = image_ins.size[0]
        orig_height = image_ins.size[1]
        image_ins = image_ins.resize((self.width, self.height))
        """
        image_ins = image_ins.convert('RGB') #convert image to RGB channel
        image_tensor = self.transform(image_ins)
    else:
        image_tensor = torch.Tensor([0.0])
        
    # Prepare target dict for torchvision model
    boxes = self.ann[["x1", "y1", "x2", "y2"]] \
                     [self.ann.id==self.im.id[index]].values.tolist()
    """
    # Resize boxes coords
    if (orig_width != self.width) or (orig_height != self.height):
        x_ratio = self.width  / orig_width
        y_ratio = self.height / orig_height
        boxes = [[b[0]*x_ratio, b[1]*y_ratio, b[2]*x_ratio, b[3]*y_ratio] for b in boxes]
    """        
    len_boxes = len(boxes)
    # Handle empty bounding boxes
    if len_boxes == 0:
        boxes = torch.zeros((0, 4), dtype=torch.float32)
    else:
        boxes = torch.as_tensor(boxes, dtype=torch.float32) 
    labels = torch.ones((len_boxes,), dtype=torch.int64)
    image_id = torch.tensor([index])
    area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])
    iscrowd = torch.zeros((len_boxes,), dtype=torch.int64)
    
    target = {}
    target["boxes"] = boxes
    target["labels"] = labels
    target["image_id"] =  image_id
    target["area"] = area
    target["iscrowd"] = iscrowd
    return image_tensor, target

# Show image with annotations
def show_data(self, index, width=None, height=None, show_annotation=True,\
              RoI_model=None, RoI_try_num=None):
    """
    Args:
        index: index of the image
        width, height: displaing size, we take default values from dataset
        show_annotaion: True - add annotation boxes
        RoI_model, RoI_try_num - optional, add detected RoI with saved model and try_num
    """
    if width==None or height==None:
        width = self.im.width[index]
        height = self.im.height[index]
    
    fig, ax = plt.subplots()
    # reshape when increase
    #ax.imshow(self[index][0].detach().numpy().reshape()) #cmap='gray'
    ax.imshow(self[index][0].detach().numpy().reshape(height, width, 3))
    ax.set_title('file = '+ str(self.im.file_name[index]))
    
    if show_annotation:
        # Select all annotation boxes for this image
        boxes_df = self.ann[self.ann.id==self.im.id[index]]
        if boxes_df.shape[0] > 0:
            x_scale = width  / self.im.width[index]
            y_scale = height / self.im.height[index]
            ann_boxes =  [Rectangle((row[0]*x_scale, row[1]*y_scale),\
                                   (row[2]-row[0])*x_scale, (row[3]-row[1])*y_scale) \
                         for row in boxes_df[["x1", "y1", "x2", "y2"]].to_numpy()]

            # Create patch collection
            pc = PatchCollection(ann_boxes, facecolor="none", edgecolor="blue", alpha=0.5)
            # Add collection to axes
            ax.add_collection(pc)
            blue_patch = mpatches.Patch(color='blue', label='Annotations from a trainer')
            red_patch = mpatches.Patch(color='red', label="Detected RoI's")
            ax.legend(handles=[blue_patch, red_patch], fontsize="small")
            
    # Model and try_num for RoI is selected to add RoI boxes
    if type(RoI_model)==str and type(RoI_try_num)==int:
        boxes_df = self.RoI[(self.RoI.im_id==self.im.iloc[index, 0]) \
                           &(self.RoI.model==RoI_model) \
                           &(self.RoI.try_num==RoI_try_num)]
        if boxes_df.shape[0] > 0:
            x_scale = width  / self.im.width[index]
            y_scale = height / self.im.height[index]
            RoI_boxes =  [Rectangle((row[0]*x_scale, row[1]*y_scale),\
                                   (row[2]-row[0])*x_scale, (row[3]-row[1])*y_scale) \
                          for row in boxes_df[["x", "y", "dx", "dy"]].to_numpy()]

            # Create patch collection
            pc = PatchCollection(RoI_boxes, facecolor="none", edgecolor="red", alpha=0.5)
            # Add collection to axes
            ax.add_collection(pc)

# Add images to dataset from directory
def add_images(self, directory="C:\\temp\\datasets\\mediag\\images\\expert_images\\"):
    """
    Args:
        directory: the directory with images to add in this dataset
    """
    
    # Select all images from directory to list
    prep_files = [f for f in os.listdir(directory) \
                                  if os.path.isfile(os.path.join(directory, f)) \
                                 and (f[-4:]==".jpg" or f[-5:]==".jpeg") \
                                 and os.stat(os.path.join(directory, f)).st_size>1000]
    # Copy file into dataset folder
    for f in prep_files:
        shutil.copy(os.path.join(directory, f), os.path.join(self.images_path, f))
    # fill parameters for new images
    prep_im = pd.DataFrame(columns = self.im.columns)
    prep_im["file_name"] = prep_files
    prep_im["cell_type"] = prep_im["file_name"].str.split("_", n=1, expand=True)[0]
    full_file_name = [os.path.join(self.images_path, f_n)
                      for f_n in prep_im["file_name"]]
    prep_im_size = pd.DataFrame([image_size(f_f_n) for f_f_n in full_file_name])
    prep_im["width"]  = prep_im_size[0]
    prep_im["height"] = prep_im_size[1]
    max_id = self.im.id.max()
    prep_im["id"] = max_id + prep_im.index + 1
    
    ## drop dublicates from dataset (detected by file_name)
    dubl_files = self.im[self.im["file_name"].isin(prep_files)].index
    self.im.drop(dubl_files, axis=0, inplace=True)
    
    self.im = pd.concat([prep_im, self.im])
    self.im.reset_index(drop=True, inplace=True)
    
    # import expert annotations (TBD)
    
    # update dataset parameters
    self.len = self.im.shape[0]
问题定位与解决

核心问题:维度转换错误

show_data方法中使用reshape(height, width, 3)调整图像张量维度是错误的:

  • 经过transform(通常包含ToTensor())处理后,image_tensor的形状为(3, height, width)(通道在前,高、宽在后)。
  • reshape仅重新排列维度的逻辑结构,不会改变数据的存储顺序,导致RGB通道的像素被混乱填充,最终显示异常。
  • 旧图像显示正常是尺寸巧合下的偶然结果,并非正确处理方式。

解决方案

将show_data中的reshape替换为transpose(1, 2, 0),该方法会正确交换维度,将(3, height, width)转换为Matplotlib所需的(height, width, 3)格式:

修改后的代码行:

ax.imshow(self[index][0].detach().numpy().transpose(1, 2, 0))

额外注意事项

如果transform中包含Normalize操作,会将像素值归一化到[-1, 1]或其他非[0,1]范围,也会导致imshow显示异常。此时需要将像素值还原到[0,1]范围:

img_np = self[index][0].detach().numpy().transpose(1, 2, 0)
# 假设Normalize的均值为mean=[0.485, 0.456, 0.406],标准差为std=[0.229, 0.224, 0.225]
img_np = img_np * std + mean
img_np = np.clip(img_np, 0, 1)
ax.imshow(img_np)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 17:54:55