diff --git a/trainer/config.py b/trainer/config.py index f02405d..c76e597 100644 --- a/trainer/config.py +++ b/trainer/config.py @@ -2,6 +2,8 @@ from typing import Union, List, Dict from pydantic import BaseModel import json from typing import Literal +# Parse the model: +from trainer.models import pretrained_models class TrainingConfig(BaseModel): output_dir: str @@ -56,6 +58,7 @@ class TrainingConfig(BaseModel): lr_power: float = 1.0 dataloader_num_workers: int = 0 training_attributes: dict = {} + aspect_ratio_bucketing: bool = True def save_as_json(self, file_path: str) -> None: with open(file_path, 'w') as f: diff --git a/trainer/dataset_and_utils.py b/trainer/dataset_and_utils.py index 9f41904..6a15c0f 100755 --- a/trainer/dataset_and_utils.py +++ b/trainer/dataset_and_utils.py @@ -111,7 +111,7 @@ def plot_loss(losses, save_path='losses.png', window_length=31, polyorder=3): def prepare_image( - pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None + pil_image: PIL.Image.Image, w: int = 512, h: int = 512, pipe=None, ) -> torch.Tensor: pil_image = pil_image.resize((w, h), resample=Image.BICUBIC, reducing_gap=1) image = pipe.image_processor.preprocess(pil_image) @@ -143,6 +143,8 @@ class PreprocessedDataset(Dataset): size: int = 512, text_dropout: float = 0.0, scale_vae_latents: bool = True, + aspect_ratio_bucketing: bool = False, + train_batch_size: int = None,# required for aspect_ratio_bucketing substitute_caption_map: Dict[str, str] = {}, ): super().__init__() @@ -203,19 +205,68 @@ class PreprocessedDataset(Dataset): else: self.do_cache = False + if aspect_ratio_bucketing: + assert train_batch_size is not None, f"Please also provide a `train_batch_size` when you have set `aspect_ratio_bucketing == True`" + from .utils.aspect_ratio_bucketing import BucketManager + aspect_ratios = {} + for idx in range(len(self.data)): + aspect_ratios[idx] = Image.open(os.path.join(os.path.dirname(self.csv_path), self.image_path[idx])).size + self.bucket_manager = BucketManager( + aspect_ratios = aspect_ratios, + bsz = train_batch_size + ) + else: + self.bucket_manager = None + + def get_aspect_ratio_bucketed_batch(self): + assert self.bucket_manager is not None, f"Expected self.bucket_manager to not be None! In order to get an aspect ration bucketed batch, please set aspect_ratio_bucketing = True and set a value for train_batch_size when doing __init__()" + indices, resolution = self.bucket_manager.get_batch() + + tok1, tok2, vae_latents, masks = [], [], [], [] + + for idx in indices: + + if self.tokenizer_2 is None: + t1, v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution) + else: + (t1, t2), v, m = self.__getitem__(idx = idx, bucketing_resolution=resolution) + tok2.append(t2.unsqueeze(0)) + + tok1.append(t1.unsqueeze(0)) + vae_latents.append(v.unsqueeze(0)) + masks.append(m.unsqueeze(0)) + + tok1 = torch.cat(tok1, dim = 0) + if self.tokenizer_2 is None: + pass + else: + tok2 = torch.cat(tok2, dim = 0) + vae_latents = torch.cat(vae_latents, dim = 0) + masks = torch.cat(masks, dim = 0) + + if self.tokenizer_2 is None: + return tok1, vae_latents, masks + else: + return (tok1, tok2), vae_latents, masks + def __len__(self) -> int: return len(self.data) @torch.no_grad() def _process( - self, idx: int + self, idx: int, bucketing_resolution: tuple = None ) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: image_path = self.image_path[idx] image_path = os.path.join(os.path.dirname(self.csv_path), image_path) image = PIL.Image.open(image_path).convert("RGB") - image = prepare_image(image, self.size, self.size, self.pipe).to( - dtype=self.vae_encoder.dtype, device=self.vae_encoder.device - ) + if bucketing_resolution is None: + image = prepare_image(image, w = self.size, h = self.size, self.pipe).to( + dtype=self.vae_encoder.dtype, device=self.vae_encoder.device + ) + else: + image = prepare_image(image, w = bucketing_resolution[0], h = bucketing_resolution[1]).to( + dtype=self.vae_encoder.dtype, device=self.vae_encoder.device + ) caption = self.caption[idx] print(caption) @@ -274,7 +325,7 @@ class PreprocessedDataset(Dataset): return (ti1, ti2), vae_latent, mask.squeeze() def __getitem__( - self, idx: int + self, idx: int, bucketing_resolution:tuple = None ) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, torch.Tensor]: if self.do_cache: vae_latent = self.vae_latents[idx].sample() @@ -282,7 +333,7 @@ class PreprocessedDataset(Dataset): vae_latent *= self.vae_scaling_factor return self.tokens_tuple[idx], vae_latent.squeeze(), self.masks[idx] else: - tokens, vae_latent, mask = self._process(idx) + tokens, vae_latent, mask = self._process(idx, bucketing_resolution=bucketing_resolution) vae_latent = vae_latent.sample() if self.scale_vae_latents: vae_latent *= self.vae_scaling_factor diff --git a/trainer/utils/aspect_ratio_bucketing.py b/trainer/utils/aspect_ratio_bucketing.py new file mode 100644 index 0000000..36691ae --- /dev/null +++ b/trainer/utils/aspect_ratio_bucketing.py @@ -0,0 +1,267 @@ +# Released under MIT license +# Copyright (c) 2022 finetuneanon (NovelAI/Anlatan LLC) + +import numpy as np +import pickle +import time + +def get_prng(seed): + return np.random.RandomState(seed) + +class BucketManager: + def __init__(self, aspect_ratios, valid_ids=None, max_size=(768,512), divisible=64, step_size=8, min_dim=256, base_res=(512,512), bsz=1, world_size=1, global_rank=0, max_ar_error=4, seed=42, dim_limit=2048, debug=False): + + self.res_map = aspect_ratios + if valid_ids is not None: + new_res_map = {} + valid_ids = set(valid_ids) + for k, v in self.res_map.items(): + if k in valid_ids: + new_res_map[k] = v + self.res_map = new_res_map + self.max_size = max_size + self.f = 8 + self.max_tokens = (max_size[0]/self.f) * (max_size[1]/self.f) + self.div = divisible + self.min_dim = min_dim + self.dim_limit = dim_limit + self.base_res = base_res + self.bsz = bsz + self.world_size = world_size + self.global_rank = global_rank + self.max_ar_error = max_ar_error + self.prng = get_prng(seed) + epoch_seed = self.prng.tomaxint() % (2**32-1) + self.epoch_prng = get_prng(epoch_seed) # separate prng for sharding use for increased thread resilience + self.epoch = None + self.left_over = None + self.batch_total = None + self.batch_delivered = None + + self.debug = debug + + self.gen_buckets() + self.assign_buckets() + self.start_epoch() + + def gen_buckets(self): + if self.debug: + timer = time.perf_counter() + resolutions = [] + aspects = [] + w = self.min_dim + while (w/self.f) * (self.min_dim/self.f) <= self.max_tokens and w <= self.dim_limit: + h = self.min_dim + got_base = False + while (w/self.f) * ((h+self.div)/self.f) <= self.max_tokens and (h+self.div) <= self.dim_limit: + if w == self.base_res[0] and h == self.base_res[1]: + got_base = True + h += self.div + if (w != self.base_res[0] or h != self.base_res[1]) and got_base: + resolutions.append(self.base_res) + aspects.append(1) + resolutions.append((w, h)) + aspects.append(float(w)/float(h)) + w += self.div + h = self.min_dim + while (h/self.f) * (self.min_dim/self.f) <= self.max_tokens and h <= self.dim_limit: + w = self.min_dim + got_base = False + while (h/self.f) * ((w+self.div)/self.f) <= self.max_tokens and (w+self.div) <= self.dim_limit: + if w == self.base_res[0] and h == self.base_res[1]: + got_base = True + w += self.div + resolutions.append((w, h)) + aspects.append(float(w)/float(h)) + h += self.div + res_map = {} + for i, res in enumerate(resolutions): + res_map[res] = aspects[i] + self.resolutions = sorted(res_map.keys(), key=lambda x: x[0] * 4096 - x[1]) + self.aspects = np.array(list(map(lambda x: res_map[x], self.resolutions))) + self.resolutions = np.array(self.resolutions) + if self.debug: + timer = time.perf_counter() - timer + print(f"resolutions:\n{self.resolutions}") + print(f"aspects:\n{self.aspects}") + print(f"gen_buckets: {timer:.5f}s") + + def assign_buckets(self): + if self.debug: + timer = time.perf_counter() + self.buckets = {} + self.aspect_errors = [] + skipped = 0 + skip_list = [] + for post_id in self.res_map.keys(): + w, h = self.res_map[post_id] + aspect = float(w)/float(h) + bucket_id = np.abs(self.aspects - aspect).argmin() + if bucket_id not in self.buckets: + self.buckets[bucket_id] = [] + error = abs(self.aspects[bucket_id] - aspect) + if error < self.max_ar_error: + self.buckets[bucket_id].append(post_id) + if self.debug: + self.aspect_errors.append(error) + else: + skipped += 1 + skip_list.append(post_id) + for post_id in skip_list: + del self.res_map[post_id] + if self.debug: + timer = time.perf_counter() - timer + self.aspect_errors = np.array(self.aspect_errors) + print(f"skipped images: {skipped}") + print(f"aspect error: mean {self.aspect_errors.mean()}, median {np.median(self.aspect_errors)}, max {self.aspect_errors.max()}") + for bucket_id in reversed(sorted(self.buckets.keys(), key=lambda b: len(self.buckets[b]))): + print(f"bucket {bucket_id}: {self.resolutions[bucket_id]}, aspect {self.aspects[bucket_id]:.5f}, entries {len(self.buckets[bucket_id])}") + print(f"assign_buckets: {timer:.5f}s") + + def start_epoch(self, world_size=None, global_rank=None): + if self.debug: + timer = time.perf_counter() + if world_size is not None: + self.world_size = world_size + if global_rank is not None: + self.global_rank = global_rank + + # select ids for this epoch/rank + index = np.array(sorted(list(self.res_map.keys()))) + index_len = index.shape[0] + index = self.epoch_prng.permutation(index) + index = index[:index_len - (index_len % (self.bsz * self.world_size))] + #print("perm", self.global_rank, index[0:16]) + index = index[self.global_rank::self.world_size] + self.batch_total = index.shape[0] // self.bsz + assert(index.shape[0] % self.bsz == 0) + index = set(index) + + self.epoch = {} + self.left_over = [] + self.batch_delivered = 0 + for bucket_id in sorted(self.buckets.keys()): + if len(self.buckets[bucket_id]) > 0: + self.epoch[bucket_id] = np.array([post_id for post_id in self.buckets[bucket_id] if post_id in index], dtype=np.int64) + self.prng.shuffle(self.epoch[bucket_id]) + self.epoch[bucket_id] = list(self.epoch[bucket_id]) + overhang = len(self.epoch[bucket_id]) % self.bsz + if overhang != 0: + self.left_over.extend(self.epoch[bucket_id][:overhang]) + self.epoch[bucket_id] = self.epoch[bucket_id][overhang:] + if len(self.epoch[bucket_id]) == 0: + del self.epoch[bucket_id] + + if self.debug: + timer = time.perf_counter() - timer + count = 0 + for bucket_id in self.epoch.keys(): + count += len(self.epoch[bucket_id]) + print(f"correct item count: {count == len(index)} ({count} of {len(index)})") + print(f"start_epoch: {timer:.5f}s") + + def get_batch(self): + if self.debug: + timer = time.perf_counter() + # check if no data left or no epoch initialized + if self.epoch is None or self.left_over is None or (len(self.left_over) == 0 and not bool(self.epoch)) or self.batch_total == self.batch_delivered: + self.start_epoch() + + found_batch = False + batch_data = None + resolution = self.base_res + while not found_batch: + bucket_ids = list(self.epoch.keys()) + if len(self.left_over) >= self.bsz: + bucket_probs = [len(self.left_over)] + [len(self.epoch[bucket_id]) for bucket_id in bucket_ids] + bucket_ids = [-1] + bucket_ids + else: + bucket_probs = [len(self.epoch[bucket_id]) for bucket_id in bucket_ids] + bucket_probs = np.array(bucket_probs, dtype=np.float32) + bucket_lens = bucket_probs + bucket_probs = bucket_probs / bucket_probs.sum() + bucket_ids = np.array(bucket_ids, dtype=np.int64) + if bool(self.epoch): + chosen_id = int(self.prng.choice(bucket_ids, 1, p=bucket_probs)[0]) + else: + chosen_id = -1 + + if chosen_id == -1: + # using leftover images that couldn't make it into a bucketed batch and returning them for use with basic square image + self.prng.shuffle(self.left_over) + batch_data = self.left_over[:self.bsz] + self.left_over = self.left_over[self.bsz:] + found_batch = True + else: + if len(self.epoch[chosen_id]) >= self.bsz: + # return bucket batch and resolution + batch_data = self.epoch[chosen_id][:self.bsz] + self.epoch[chosen_id] = self.epoch[chosen_id][self.bsz:] + resolution = tuple(self.resolutions[chosen_id]) + found_batch = True + if len(self.epoch[chosen_id]) == 0: + del self.epoch[chosen_id] + else: + # can't make a batch from this, not enough images. move them to leftovers and try again + self.left_over.extend(self.epoch[chosen_id]) + del self.epoch[chosen_id] + + assert(found_batch or len(self.left_over) >= self.bsz or bool(self.epoch)) + + if self.debug: + timer = time.perf_counter() - timer + print(f"bucket probs: " + ", ".join(map(lambda x: f"{x:.2f}", list(bucket_probs*100)))) + print(f"chosen id: {chosen_id}") + print(f"batch data: {batch_data}") + print(f"resolution: {resolution}") + print(f"get_batch: {timer:.5f}s") + + self.batch_delivered += 1 + return (batch_data, resolution) + + def generator(self): + if self.batch_delivered >= self.batch_total: + self.start_epoch() + while self.batch_delivered < self.batch_total: + yield self.get_batch() + +if __name__ == "__main__": + # prepare a pickle with mapping of dataset IDs to resolutions called resolutions.pkl to use this + with open("resolutions.pkl", "rb") as fh: + ids = list(pickle.load(fh).keys()) + + counts = np.zeros((len(ids),)).astype(np.int64) + id_map = {} + for i, post_id in enumerate(ids): + id_map[post_id] = i + + bm = BucketManager("resolutions.pkl", debug=True, bsz=8, world_size=8, global_rank=3) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + print("got: " + str(bm.get_batch())) + + bm = BucketManager("resolutions.pkl", bsz=8, world_size=1, global_rank=0, valid_ids=ids[0:16]) + for _ in range(16): + bm.get_batch() + print("got from future epoch: " + str(bm.get_batch())) + + bms = [] + for rank in range(16): + bm = BucketManager("resolutions.pkl", bsz=8, world_size=16, global_rank=rank) + bms.append(bm) + for epoch in range(5): + print(f"epoch {epoch}") + for i, bm in enumerate(bms): + print(f"bm {i}") + first = True + for ids, res in bm.generator(): + if first and i == 0: + #print(ids) + first = False + for post_id in ids: + counts[id_map[post_id]] += 1 + print(np.bincount(counts)) \ No newline at end of file diff --git a/trainer_pti.py b/trainer_pti.py index 6fc4edf..20f5eca 100755 --- a/trainer_pti.py +++ b/trainer_pti.py @@ -237,6 +237,8 @@ def main( vae, do_cache=config.do_cache, substitute_caption_map=config.token_dict, + aspect_ratio_bucketing=config.train_batch_size, + train_batch_size=config.train_batch_size ) print(f"# PTI : Loaded dataset, do_cache: {config.do_cache}") @@ -301,6 +303,8 @@ def main( start_time, images_done = time.time(), 0 for epoch in range(first_epoch, config.num_train_epochs): + if config.aspect_ratio_bucketing: + train_dataset.bucket_manager.start_epoch() progress_bar.set_description(f"# PTI :step: {global_step}, epoch: {epoch}") for step, batch in enumerate(train_dataloader): @@ -327,10 +331,21 @@ def main( if config.pretrained_model['version'] == "sdxl": - (tok1, tok2), vae_latent, mask = batch + if not config.aspect_ratio_bucketing: + (tok1, tok2), vae_latent, mask = batch + else: + ## delete batch to save a bit of memory + del batch + (tok1, tok2), vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch() elif config.pretrained_model['version'] == "sd15": - tok1, vae_latent, mask = batch + if not config.aspect_ratio_bucketing: + tok1, vae_latent, mask = batch + else: + ## delete batch to save a bit of memory + del batch + tok1, vae_latent, mask = train_dataset.get_aspect_ratio_bucketed_batch() tok2 = None + vae_latent = vae_latent.to(weight_dtype) diff --git a/training_args.json b/training_args.json index cef2e1b..d0ad775 100644 --- a/training_args.json +++ b/training_args.json @@ -27,5 +27,6 @@ "debug": true, "hard_pivot": false, "mixed_precision": "bf16", - "dataloader_num_workers": 0 + "dataloader_num_workers": 0, + "aspect_ratio_bucketing": true } \ No newline at end of file