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

如何在PyTorch中实现可采样不同风格图像的Dataset/Dataloader?

风格迁移任务的PyTorch Dataset实现方案

先指出你初步实现里的几个问题:

  1. random.shuffle()是原地修改列表,返回值为None,你的代码会把itr_style和itr_content赋值为None,后续调用直接报错。
  2. __init__仅在Dataset初始化时执行一次,无法实现每个epoch重新洗牌配对的需求。
  3. 没有风格校验逻辑,可能出现风格图和内容图属于同一风格的情况。

下面是满足你需求的最优实现方案:

核心思路

  1. 初始化时按风格对图像索引做分组,避免每次生成配对时重复遍历所有图像。
  2. 单独实现配对生成方法,确保风格图和内容图来自不同风格。
  3. 提供on_epoch_end()方法,在每个epoch开始时重新生成配对,实现全epoch洗牌。

完整代码实现

import random
from torch.utils.data import Dataset

class StyleTransferDataset(Dataset):
    def __init__(self, images, style_labels):
        # images: 图像列表(可以是图像路径或已加载的张量,根据实际场景调整)
        # style_labels: 对应每个图像的风格标签列表(如字符串/整数)
        self.images = images
        self.style_labels = style_labels
        self.num_samples = len(images)
        
        # 按风格分组图像索引,初始化时仅执行一次
        self.style_groups = {}
        for idx, label in enumerate(style_labels):
            if label not in self.style_groups:
                self.style_groups[label] = []
            self.style_groups[label].append(idx)
        self.all_styles = list(self.style_groups.keys())
        
        # 生成初始配对
        self.pairs = self._generate_pairs()
    
    def _generate_pairs(self):
        """生成满足风格不同要求的图像配对"""
        pairs = []
        # 先打乱内容图的顺序,保证每个epoch内容图的顺序不同
        content_indices = list(range(self.num_samples))
        random.shuffle(content_indices)
        
        for content_idx in content_indices:
            content_style = self.style_labels[content_idx]
            # 筛选出当前内容图风格之外的所有风格
            available_styles = [s for s in self.all_styles if s != content_style]
            # 随机选一个目标风格
            target_style = random.choice(available_styles)
            # 从目标风格组里随机选一张图作为风格图
            style_idx = random.choice(self.style_groups[target_style])
            pairs.append((content_idx, style_idx))
        return pairs
    
    def on_epoch_end(self):
        """每个epoch结束/开始时调用,重新生成配对"""
        self.pairs = self._generate_pairs()
    
    def __len__(self):
        return self.num_samples
    
    def __getitem__(self, idx):
        content_idx, style_idx = self.pairs[idx]
        # 如果是图像路径,这里需要添加图像读取逻辑(如PIL/OpenCV)
        # 示例:content_img = Image.open(self.images[content_idx]).convert('RGB')
        content_img = self.images[content_idx]
        style_img = self.images[style_idx]
        # 可在此处添加transform操作
        return content_img, style_img

使用方式

训练时需要在每个epoch开始前调用on_epoch_end()方法,确保配对重新洗牌:

from torch.utils.data import DataLoader

# 假设已准备好images(图像路径/张量列表)和style_labels(对应风格标签)
dataset = StyleTransferDataset(images, style_labels)
# 此处shuffle设为False,因为我们通过on_epoch_end()自行管理配对洗牌
dataloader = DataLoader(dataset, batch_size=8, shuffle=False)

num_epochs = 10
for epoch in range(num_epochs):
    # 每个epoch开始前重新生成配对
    dataset.on_epoch_end()
    for content_batch, style_batch in dataloader:
        # 你的训练逻辑
        pass

方案优势

  • 高效性:风格分组仅在初始化时执行一次,后续生成配对直接复用分组结果。
  • 合规性:严格保证风格图与内容图来自不同风格。
  • 随机性:每个epoch通过重新生成配对实现全样本洗牌,避免模型过拟合固定配对。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 12:39:50