torchvision中如何为不同标签的图像应用差异化图像变换?
按标签/索引为不同类别图像添加颜色方块的解决方案
一、按标签为图像添加对应颜色方块
ImageFolder返回的每个样本是(图像, 标签)对,但默认的transform只作用于图像。要根据标签做变换,最直接的方式是自定义一个继承自ImageFolder的数据集类,在__getitem__中手动处理图像和标签的映射:
完整代码示例
import numpy as np from PIL import Image from torchvision import datasets, transforms import os # 定义颜色与标签的映射:标签0→红,1→蓝,2→绿 LABEL_COLOR_MAP = { 0: [255, 0, 0], 1: [0, 0, 255], 2: [0, 255, 0] } class LabelColoredSquareDataset(datasets.ImageFolder): def __getitem__(self, index): # 获取原始图像和标签 img, label = super().__getitem__(index) # 将PIL图像转为numpy数组,添加颜色方块 img_np = np.array(img) # 左上角20x20区域填充对应颜色 img_np[0:20, 0:20, :] = LABEL_COLOR_MAP[label] # 转回PIL图像,应用基础变换(如转Tensor) img = Image.fromarray(img_np) if self.transform is not None: img = self.transform(img) return img, label # 定义基础变换(替换成你原来的transforms.Compose) base_transform = transforms.Compose([ transforms.ToTensor() ]) # 初始化数据集 data_dir = "你的数据路径" perturb_dataset = LabelColoredSquareDataset( root=os.path.join(data_dir, 'training'), transform=base_transform ) # 创建DataLoader batch_size = 32 perturb_loader = torch.utils.data.DataLoader( perturb_dataset, batch_size=batch_size, shuffle=True )
说明
- 自定义数据集类继承自
ImageFolder,复用了它自动读取文件夹标签的逻辑 - 在
__getitem__中先获取标签,再根据标签映射的颜色修改图像像素 - 最后再应用你原本的基础变换(如转Tensor)
二、无法获取标签时,按图像索引实现颜色区分
如果因为某些原因无法直接获取标签,也可以通过图像的索引来分配颜色。常见的两种方式:
方式1:按索引取模分配颜色(适合随机分配)
这种方式不依赖样本的实际类别,仅根据索引的模值分配颜色,适合不需要和类别严格对应的场景:
class IndexColoredSquareDataset(datasets.ImageFolder): def __getitem__(self, index): img, _ = super().__getitem__(index) img_np = np.array(img) # 索引模3:0→红,1→蓝,2→绿 color_idx = index % 3 color_map = {0: [255,0,0], 1: [0,0,255], 2: [0,255,0]} img_np[0:20, 0:20, :] = color_map[color_idx] img = Image.fromarray(img_np) if self.transform is not None: img = self.transform(img) return img, _
方式2:按样本区间分配颜色(模拟标签对应)
如果知道每个类别样本的数量,可以根据索引所在的区间来分配颜色,实现和标签对应的效果:
class RangeBasedColoredSquareDataset(datasets.ImageFolder): def __init__(self, root, transform=None): super().__init__(root, transform) # 计算每个类别的样本数量,生成区间边界 class_dirs = self.classes class_counts = [] for cls in class_dirs: cls_path = os.path.join(root, cls) class_counts.append(len([f for f in os.listdir(cls_path) if os.path.isfile(os.path.join(cls_path, f))])) # 计算累计样本数,作为区间分割点 self.cumulative_counts = np.cumsum(class_counts) def __getitem__(self, index): img, _ = super().__getitem__(index) img_np = np.array(img) # 根据索引所在区间分配颜色 if index < self.cumulative_counts[0]: color = [255, 0, 0] # 第一类:红 elif index < self.cumulative_counts[1]: color = [0, 0, 255] # 第二类:蓝 else: color = [0, 255, 0] # 第三类:绿 img_np[0:20, 0:20, :] = color img = Image.fromarray(img_np) if self.transform is not None: img = self.transform(img) return img, _
说明
- 方式1的颜色分配随索引变化,如果开启
shuffle=True,同一图像在不同epoch可能被分配不同颜色 - 方式2依赖样本的原始顺序(
ImageFolder会按文件夹顺序读取样本),如果不开启shuffle,颜色分配会和实际类别严格对应
内容的提问来源于stack exchange,提问作者Steve
相关产品推荐
相关产品推荐

