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

