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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 21:37:48