import base64 import glob import hashlib import itertools import json import math import os import random from collections import OrderedDict, deque from concurrent.futures import ThreadPoolExecutor from typing import TYPE_CHECKING, List, Dict, Union import traceback import cv2 import numpy as np import torch from safetensors.torch import load_file, save_file from tqdm import tqdm from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection, SiglipImageProcessor from toolkit.audio.preserve_pitch import time_stretch_preserve_pitch from toolkit.basic import flush, value_map from toolkit.buckets import get_bucket_for_image_size, get_resolution from toolkit.config_modules import ControlTypes from toolkit.control_generator import ControlGenerator from toolkit.dto import DTO, DISK_PREFIX from toolkit.metadata import get_meta_for_safetensors from toolkit.models.pixtral_vision import PixtralVisionImagePreprocessorCompatible from toolkit.prompt_utils import inject_trigger_into_prompt from torchvision import transforms from PIL import Image, ImageFilter, ImageOps from PIL.ImageOps import exif_transpose import albumentations as A from toolkit.print import print_acc from toolkit.accelerator import get_accelerator from toolkit.prompt_utils import PromptEmbeds from torchvision.transforms import functional as TF from toolkit.train_tools import get_torch_dtype if TYPE_CHECKING: from toolkit.data_loader import AiToolkitDataset from toolkit.data_transfer_object.data_loader import FileItemDTO from toolkit.stable_diffusion_model import StableDiffusion accelerator = get_accelerator() # def get_associated_caption_from_img_path(img_path): # https://demo.albumentations.ai/ class Augments: def __init__(self, **kwargs): self.method_name = kwargs.get('method', None) self.params = kwargs.get('params', {}) # convert kwargs enums for cv2 for key, value in self.params.items(): if isinstance(value, str): # split the string split_string = value.split('.') if len(split_string) == 2 and split_string[0] == 'cv2': if hasattr(cv2, split_string[1]): self.params[key] = getattr(cv2, split_string[1].upper()) else: raise ValueError(f"invalid cv2 enum: {split_string[1]}") transforms_dict = { 'ColorJitter': transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.03), 'RandomEqualize': transforms.RandomEqualize(p=0.2), } img_ext_list = ['.jpg', '.jpeg', '.png', '.webp'] video_ext_list = ['.mp4', '.avi', '.mov', '.webm', '.mkv', '.wmv', '.m4v', '.flv'] def standardize_images(images): """ Standardize the given batch of images using the specified mean and std. Expects values of 0 - 1 Args: images (torch.Tensor): A batch of images in the shape of (N, C, H, W), where N is the number of images, C is the number of channels, H is the height, and W is the width. Returns: torch.Tensor: Standardized images. """ mean = [0.48145466, 0.4578275, 0.40821073] std = [0.26862954, 0.26130258, 0.27577711] # Define the normalization transform normalize = transforms.Normalize(mean=mean, std=std) # Apply normalization to each image in the batch standardized_images = torch.stack([normalize(img) for img in images]) return standardized_images def clean_caption(caption): # this doesnt make any sense anymore in a world that is not based on comma seperated tokens # # remove any newlines # caption = caption.replace('\n', ', ') # # remove new lines for all operating systems # caption = caption.replace('\r', ', ') # caption_split = caption.split(',') # # remove empty strings # caption_split = [p.strip() for p in caption_split if p.strip()] # # join back together # caption = ', '.join(caption_split) return caption def waveform_to_stereo(waveform): c = waveform.shape[0] if c == 2: return waveform if c == 1: return waveform.expand(2, -1) if c == 6: # 5.1: FL, FR, FC, LFE, BL, BR fl, fr, fc, _, bl, br = waveform k = 0.7071 return torch.stack([fl + k * fc + k * bl, fr + k * fc + k * br]) if c == 8: # 7.1: FL, FR, FC, LFE, BL, BR, SL, SR fl, fr, fc, _, bl, br, sl, sr = waveform k = 0.7071 return torch.stack([fl + k * fc + k * (bl + sl), fr + k * fc + k * (br + sr)]) return waveform.mean(0, keepdim=True).expand(2, -1) class CaptionMixin: def get_caption_item(self: 'AiToolkitDataset', index): if not hasattr(self, 'caption_type'): raise Exception('caption_type not found on class instance') if not hasattr(self, 'file_list'): raise Exception('file_list not found on class instance') img_path_or_tuple = self.file_list[index] ext = self.dataset_config.caption_ext if isinstance(img_path_or_tuple, tuple): img_path = img_path_or_tuple[0] if isinstance(img_path_or_tuple[0], str) else img_path_or_tuple[0].path # check if either has a prompt file path_no_ext = os.path.splitext(img_path)[0] prompt_path = None prompt_path = path_no_ext + ext else: img_path = img_path_or_tuple if isinstance(img_path_or_tuple, str) else img_path_or_tuple.path # see if prompt file exists path_no_ext = os.path.splitext(img_path)[0] prompt_path = path_no_ext + ext # allow folders to have a default prompt default_prompt_path = os.path.join(os.path.dirname(img_path), 'default.txt') default_prompt_path_with_ext = os.path.join(os.path.dirname(img_path), 'default' + ext) if os.path.exists(prompt_path): with open(prompt_path, 'r', encoding='utf-8') as f: prompt = f.read() prompt = clean_caption(prompt) elif os.path.exists(default_prompt_path_with_ext): with open(default_prompt_path_with_ext, 'r', encoding='utf-8') as f: prompt = f.read() prompt = clean_caption(prompt) elif os.path.exists(default_prompt_path): with open(default_prompt_path, 'r', encoding='utf-8') as f: prompt = f.read() prompt = clean_caption(prompt) else: prompt = '' # get default_prompt if it exists on the class instance if hasattr(self, 'default_prompt'): prompt = self.default_prompt if hasattr(self, 'default_caption'): prompt = self.default_caption # handle replacements replacement_list = self.dataset_config.replacements if isinstance(self.dataset_config.replacements, list) else [] for replacement in replacement_list: from_string, to_string = replacement.split('|') prompt = prompt.replace(from_string, to_string) return prompt if TYPE_CHECKING: from toolkit.config_modules import DatasetConfig from toolkit.data_transfer_object.data_loader import FileItemDTO class Bucket: def __init__(self, width: int, height: int): self.width = width self.height = height self.file_list_idx: List[int] = [] class BucketsMixin: def __init__(self): self.buckets: Dict[str, Bucket] = {} self.batch_indices: List[List[int]] = [] def build_batch_indices(self: 'AiToolkitDataset'): self.batch_indices = [] for key, bucket in self.buckets.items(): for start_idx in range(0, len(bucket.file_list_idx), self.batch_size): end_idx = min(start_idx + self.batch_size, len(bucket.file_list_idx)) batch = bucket.file_list_idx[start_idx:end_idx] # if the bucket has fewer items left than the requested batch size, # duplicate items from this batch to pad it up to batch_size if len(batch) < self.batch_size and len(batch) > 0: pad = [batch[i % len(batch)] for i in range(self.batch_size - len(batch))] batch = batch + pad self.batch_indices.append(batch) def shuffle_buckets(self: 'AiToolkitDataset'): for key, bucket in self.buckets.items(): random.shuffle(bucket.file_list_idx) def setup_buckets(self: 'AiToolkitDataset', quiet=False): if not hasattr(self, 'file_list'): raise Exception(f'file_list not found on class instance {self.__class__.__name__}') if not hasattr(self, 'dataset_config'): raise Exception(f'dataset_config not found on class instance {self.__class__.__name__}') if self.epoch_num > 0: # no need to rebuild buckets for now # todo handle random cropping for buckets return self.buckets = {} # clear it config: 'DatasetConfig' = self.dataset_config resolution = config.resolution bucket_tolerance = config.bucket_tolerance file_list: List['FileItemDTO'] = self.file_list # for file_item in enumerate(file_list): for idx, file_item in enumerate(file_list): file_item: 'FileItemDTO' = file_item if self.is_audio_model: bucket_key = f"{file_item.width}ms" if bucket_key not in self.buckets: self.buckets[bucket_key] = Bucket(file_item.width, 1) self.buckets[bucket_key].file_list_idx.append(idx) continue width = int(file_item.width * file_item.dataset_config.scale) height = int(file_item.height * file_item.dataset_config.scale) if self.dataset_config.square_crop: # we scale first so smallest size matches resolution scale_factor_x = resolution / width scale_factor_y = resolution / height scale_factor = max(scale_factor_x, scale_factor_y) file_item.scale_to_width = math.ceil(width * scale_factor) file_item.scale_to_height = math.ceil(height * scale_factor) file_item.crop_width = resolution file_item.crop_height = resolution if width > height: file_item.crop_x = int(file_item.scale_to_width / 2 - resolution / 2) file_item.crop_y = 0 else: file_item.crop_x = 0 file_item.crop_y = int(file_item.scale_to_height / 2 - resolution / 2) else: bucket_resolution = get_bucket_for_image_size( width, height, resolution=resolution, divisibility=bucket_tolerance ) # Calculate scale factors for width and height width_scale_factor = bucket_resolution["width"] / width height_scale_factor = bucket_resolution["height"] / height # Use the maximum of the scale factors to ensure both dimensions are scaled above the bucket resolution max_scale_factor = max(width_scale_factor, height_scale_factor) # round up file_item.scale_to_width = int(math.ceil(width * max_scale_factor)) file_item.scale_to_height = int(math.ceil(height * max_scale_factor)) file_item.crop_height = bucket_resolution["height"] file_item.crop_width = bucket_resolution["width"] new_width = bucket_resolution["width"] new_height = bucket_resolution["height"] if self.dataset_config.random_crop: # random crop crop_x = random.randint(0, file_item.scale_to_width - new_width) crop_y = random.randint(0, file_item.scale_to_height - new_height) file_item.crop_x = crop_x file_item.crop_y = crop_y else: # do central crop file_item.crop_x = int((file_item.scale_to_width - new_width) / 2) file_item.crop_y = int((file_item.scale_to_height - new_height) / 2) if file_item.crop_y < 0 or file_item.crop_x < 0: print_acc('debug') # check if bucket exists, if not, create it bucket_key = f'{file_item.crop_width}x{file_item.crop_height}' if self.is_video: # images (1 frame) and videos must not mix in a batch bucket_key += f'x{file_item.num_frames}f' if bucket_key not in self.buckets: self.buckets[bucket_key] = Bucket(file_item.crop_width, file_item.crop_height) self.buckets[bucket_key].file_list_idx.append(idx) # print the buckets self.shuffle_buckets() self.build_batch_indices() if not quiet: print_acc(f'Bucket sizes for {self.dataset_path}:') for key, bucket in self.buckets.items(): print_acc(f'{key}: {len(bucket.file_list_idx)} files') print_acc(f'{len(self.buckets)} buckets made') class CaptionProcessingDTOMixin: def __init__(self: 'FileItemDTO', *args, **kwargs): if hasattr(super(), '__init__'): super().__init__(*args, **kwargs) self.raw_caption: str = None self.raw_caption_short: str = None self.caption: str = None self.caption_short: str = None # caption with the trigger word replaced by the diff output preservation class self.caption_dop: str = None # D-OPSD teacher caption (trigger word replaced by /