CrucibleAI/ControlNetMediaPipeFace
5761.2k
1import json2import numpy3import os4from PIL import Image5from torch.utils.data import Dataset6 7 8class LaionDataset(Dataset):9 def __init__(self):10 self.data = []11 with open('./training/laion-face-processed/prompt.jsonl', 'rt') as f:12 for line in f:13 self.data.append(json.loads(line))14 15 def __len__(self):16 return len(self.data)17 18 def __getitem__(self, idx):19 item = self.data[idx]20 21 source_filename = os.path.split(item['source'])[-1]22 target_filename = os.path.split(item['target'])[-1]23 prompt = item['prompt']24 25 # If prompt is "" or null, make it something simple.26 if not prompt:27 print(f"Image with index {idx} / {source_filename} has no text.")28 prompt = "an image"29 30 source_image = Image.open('./training/laion-face-processed/source/' + source_filename).convert("RGB")31 target_image = Image.open('./training/laion-face-processed/target/' + target_filename).convert("RGB")32 # Resize the image so that the minimum edge is bigger than 512x512, then crop center.33 # This may cut off some parts of the face image, but in general they're smaller than 512x512 and we still want34 # to cover the literal edge cases.35 img_size = source_image.size36 scale_factor = 512/min(img_size)37 source_image = source_image.resize((1+int(img_size[0]*scale_factor), 1+int(img_size[1]*scale_factor)))38 target_image = target_image.resize((1+int(img_size[0]*scale_factor), 1+int(img_size[1]*scale_factor)))39 img_size = source_image.size40 left_padding = (img_size[0] - 512)//241 top_padding = (img_size[1] - 512)//242 source_image = source_image.crop((left_padding, top_padding, left_padding+512, top_padding+512))43 target_image = target_image.crop((left_padding, top_padding, left_padding+512, top_padding+512))44 45 source = numpy.asarray(source_image)46 target = numpy.asarray(target_image)47 48 # Normalize source images to [0, 1].49 source = source.astype(numpy.float32) / 255.050 51 # Normalize target images to [-1, 1].52 target = (target.astype(numpy.float32) / 127.5) - 1.053 54 return dict(jpg=target, txt=prompt, hint=source)55 56 