imeplement aspect ratio bucketing
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
+17
-2
@@ -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)
|
||||
|
||||
|
||||
+2
-1
@@ -27,5 +27,6 @@
|
||||
"debug": true,
|
||||
"hard_pivot": false,
|
||||
"mixed_precision": "bf16",
|
||||
"dataloader_num_workers": 0
|
||||
"dataloader_num_workers": 0,
|
||||
"aspect_ratio_bucketing": true
|
||||
}
|
||||
Reference in New Issue
Block a user