TwoPerCent/instruct-pix2pix
0
1from __future__ import annotations2 3import json4import math5from pathlib import Path6from typing import Any7 8import numpy as np9import torch10import torchvision11from einops import rearrange12from PIL import Image13from torch.utils.data import Dataset14 15 16class EditDataset(Dataset):17 def __init__(18 self,19 path: str,20 split: str = "train",21 splits: tuple[float, float, float] = (0.9, 0.05, 0.05),22 min_resize_res: int = 256,23 max_resize_res: int = 256,24 crop_res: int = 256,25 flip_prob: float = 0.0,26 ):27 assert split in ("train", "val", "test")28 assert sum(splits) == 129 self.path = path30 self.min_resize_res = min_resize_res31 self.max_resize_res = max_resize_res32 self.crop_res = crop_res33 self.flip_prob = flip_prob34 35 with open(Path(self.path, "seeds.json")) as f:36 self.seeds = json.load(f)37 38 split_0, split_1 = {39 "train": (0.0, splits[0]),40 "val": (splits[0], splits[0] + splits[1]),41 "test": (splits[0] + splits[1], 1.0),42 }[split]43 44 idx_0 = math.floor(split_0 * len(self.seeds))45 idx_1 = math.floor(split_1 * len(self.seeds))46 self.seeds = self.seeds[idx_0:idx_1]47 48 def __len__(self) -> int:49 return len(self.seeds)50 51 def __getitem__(self, i: int) -> dict[str, Any]:52 name, seeds = self.seeds[i]53 propt_dir = Path(self.path, name)54 seed = seeds[torch.randint(0, len(seeds), ()).item()]55 with open(propt_dir.joinpath("prompt.json")) as fp:56 prompt = json.load(fp)["edit"]57 58 image_0 = Image.open(propt_dir.joinpath(f"{seed}_0.jpg"))59 image_1 = Image.open(propt_dir.joinpath(f"{seed}_1.jpg"))60 61 reize_res = torch.randint(self.min_resize_res, self.max_resize_res + 1, ()).item()62 image_0 = image_0.resize((reize_res, reize_res), Image.Resampling.LANCZOS)63 image_1 = image_1.resize((reize_res, reize_res), Image.Resampling.LANCZOS)64 65 image_0 = rearrange(2 * torch.tensor(np.array(image_0)).float() / 255 - 1, "h w c -> c h w")66 image_1 = rearrange(2 * torch.tensor(np.array(image_1)).float() / 255 - 1, "h w c -> c h w")67 68 crop = torchvision.transforms.RandomCrop(self.crop_res)69 flip = torchvision.transforms.RandomHorizontalFlip(float(self.flip_prob))70 image_0, image_1 = flip(crop(torch.cat((image_0, image_1)))).chunk(2)71 72 return dict(edited=image_1, edit=dict(c_concat=image_0, c_crossattn=prompt))73 74 75class EditDatasetEval(Dataset):76 def __init__(77 self,78 path: str,79 split: str = "train",80 splits: tuple[float, float, float] = (0.9, 0.05, 0.05),81 res: int = 256,82 ):83 assert split in ("train", "val", "test")84 assert sum(splits) == 185 self.path = path86 self.res = res87 88 with open(Path(self.path, "seeds.json")) as f:89 self.seeds = json.load(f)90 91 split_0, split_1 = {92 "train": (0.0, splits[0]),93 "val": (splits[0], splits[0] + splits[1]),94 "test": (splits[0] + splits[1], 1.0),95 }[split]96 97 idx_0 = math.floor(split_0 * len(self.seeds))98 idx_1 = math.floor(split_1 * len(self.seeds))99 self.seeds = self.seeds[idx_0:idx_1]100 101 def __len__(self) -> int:102 return len(self.seeds)103 104 def __getitem__(self, i: int) -> dict[str, Any]:105 name, seeds = self.seeds[i]106 propt_dir = Path(self.path, name)107 seed = seeds[torch.randint(0, len(seeds), ()).item()]108 with open(propt_dir.joinpath("prompt.json")) as fp:109 prompt = json.load(fp)110 edit = prompt["edit"]111 input_prompt = prompt["input"]112 output_prompt = prompt["output"]113 114 image_0 = Image.open(propt_dir.joinpath(f"{seed}_0.jpg"))115 116 reize_res = torch.randint(self.res, self.res + 1, ()).item()117 image_0 = image_0.resize((reize_res, reize_res), Image.Resampling.LANCZOS)118 119 image_0 = rearrange(2 * torch.tensor(np.array(image_0)).float() / 255 - 1, "h w c -> c h w")120 121 return dict(image_0=image_0, input_prompt=input_prompt, edit=edit, output_prompt=output_prompt)122 