如何获取PyTorch自定义分割数据集中指定索引的图像名称?
获取指定索引的图像名称
有两种简单的方法可以实现需求:
方法1:直接从数据集的DataFrame中读取
你的SegmentationDataset类保存了传入的DataFrame(self.df),可以直接通过索引从DataFrame中提取图像名称:
idx = 9 # 获取包含DATA_DIR前缀的图像路径 image_full_path = validset.df.iloc[idx]['images'] # 若只需文件名(去掉路径前缀),可使用os.path.basename import os image_name = os.path.basename(image_full_path) print(f"索引{idx}对应的图像名称:{image_name}")
方法2:修改Dataset类,支持获取图像名称
如果希望更规范,可以通过修改Dataset类实现,有两种子方式:
方式A:修改__getitem__返回额外信息
调整__getitem__方法,让它同时返回图像名称:
class SegmentationDataset(Dataset): def __init__(self, df, augmentations): self.df = df self.augmentations = augmentations def __len__(self): return len(self.df) def __getitem__(self, idx): row = self.df.iloc[idx] image_path = DATA_DIR + row.images mask_path = DATA_DIR + row.masks image = skimage.io.imread(image_path) mask = skimage.io.imread(mask_path) mask = np.expand_dims(mask, axis = -1) if self.augmentations: data = self.augmentations(image = image, mask = mask) image = data['image'] mask = data['mask'] image = np.transpose(image, (2, 0, 1)).astype(np.float32) mask = np.transpose(mask, (2, 0, 1)).astype(np.float32) image = torch.Tensor(image) / 255.0 mask = torch.round(torch.Tensor(mask) / 255.0) # 新增返回图像路径/名称 return image, mask, row.images # 使用示例 idx = 9 image, mask, image_full_path = validset[idx] image_name = os.path.basename(image_full_path) print(f"索引{idx}对应的图像名称:{image_name}")
方式B:新增专门的获取方法
如果不想改变原有__getitem__的返回值,可以给Dataset类添加一个独立方法:
class SegmentationDataset(Dataset): def __init__(self, df, augmentations): self.df = df self.augmentations = augmentations def __len__(self): return len(self.df) def __getitem__(self, idx): # 原有代码保持不变 row = self.df.iloc[idx] image_path = DATA_DIR + row.images mask_path = DATA_DIR + row.masks image = skimage.io.imread(image_path) mask = skimage.io.imread(mask_path) mask = np.expand_dims(mask, axis = -1) if self.augmentations: data = self.augmentations(image = image, mask = mask) image = data['image'] mask = data['mask'] image = np.transpose(image, (2, 0, 1)).astype(np.float32) mask = np.transpose(mask, (2, 0, 1)).astype(np.float32) image = torch.Tensor(image) / 255.0 mask = torch.round(torch.Tensor(mask) / 255.0) return image, mask # 新增获取图像名称的方法 def get_image_name(self, idx): row = self.df.iloc[idx] return row.images # 使用示例 idx = 9 image_full_path = validset.get_image_name(idx) image_name = os.path.basename(image_full_path) print(f"索引{idx}对应的图像名称:{image_name}")
内容的提问来源于stack exchange,提问作者Milap
相关产品推荐
相关产品推荐

