Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
simple_text_image.py36 linesDownload Raw Back to data
1import torch, os2from torchvision import transforms3import pandas as pd4from PIL import Image5 6 7 8class TextImageDataset(torch.utils.data.Dataset):9    def __init__(self, dataset_path, steps_per_epoch=10000, height=1024, width=1024, center_crop=True, random_flip=False):10        self.steps_per_epoch = steps_per_epoch11        metadata = pd.read_csv(os.path.join(dataset_path, "train/metadata.csv"))12        self.path = [os.path.join(dataset_path, "train", file_name) for file_name in metadata["file_name"]]13        self.text = metadata["text"].to_list()14        self.image_processor = transforms.Compose(15            [16                transforms.Resize(max(height, width), interpolation=transforms.InterpolationMode.BILINEAR),17                transforms.CenterCrop((height, width)) if center_crop else transforms.RandomCrop((height, width)),18                transforms.RandomHorizontalFlip() if random_flip else transforms.Lambda(lambda x: x),19                transforms.ToTensor(),20                transforms.Normalize([0.5], [0.5]),21            ]22        )23 24 25    def __getitem__(self, index):26        data_id = torch.randint(0, len(self.path), (1,))[0]27        data_id = (data_id + index) % len(self.path) # For fixed seed.28        text = self.text[data_id]29        image = Image.open(self.path[data_id]).convert("RGB")30        image = self.image_processor(image)31        return {"text": text, "image": image}32 33 34    def __len__(self):35        return self.steps_per_epoch36