| 396 | |
| 397 | |
| 398 | class ARBucketDataset: |
| 399 | def __init__(self, ar_frames, resolutions, metadata_dataset, directory_config, cache_base, round_to_multiple, directory_dataset): |
| 400 | self.ar_frames = ar_frames |
| 401 | self.resolutions = resolutions |
| 402 | self.metadata_dataset = metadata_dataset |
| 403 | self.directory_config = directory_config |
| 404 | self.size_buckets = [] |
| 405 | self.path = Path(directory_config['path']) |
| 406 | self.cache_base = cache_base |
| 407 | self.cache_dir = cache_base / f'ar_frames_{bucket_suffix(self.ar_frames)}' |
| 408 | self.round_to_multiple = round_to_multiple |
| 409 | self.directory_dataset = directory_dataset |
| 410 | os.makedirs(self.cache_dir, exist_ok=True) |
| 411 | |
| 412 | def get_size_bucket_datasets(self): |
| 413 | return self.size_buckets |
| 414 | |
| 415 | def cache_latents(self, map_fn, regenerate_cache=False, trust_cache=False, caching_batch_size=1): |
| 416 | print(f'caching latents: {self.ar_frames}') |
| 417 | |
| 418 | for res in self.resolutions: |
| 419 | area = res**2 |
| 420 | w = math.sqrt(area * self.ar_frames[0]) |
| 421 | h = area / w |
| 422 | w = round_to_nearest_multiple(w, self.round_to_multiple) |
| 423 | h = round_to_nearest_multiple(h, self.round_to_multiple) |
| 424 | size_bucket = (w, h, self.ar_frames[1]) |
| 425 | # to make sure the directory has a unique name |
| 426 | naming_size_bucket = (self.ar_frames[0],) + size_bucket |
| 427 | metadata_with_size_bucket = self.metadata_dataset.map( |
| 428 | lambda example: {'size_bucket': size_bucket}, |
| 429 | cache_file_name=str(self.cache_dir / f'metadata/metadata_{bucket_suffix(naming_size_bucket)}.arrow'), |
| 430 | load_from_cache_file=(not regenerate_cache and trust_cache), |
| 431 | desc='Adding size bucket', |
| 432 | ) |
| 433 | self.size_buckets.append( |
| 434 | SizeBucketDataset(metadata_with_size_bucket, self.directory_config, naming_size_bucket, self.cache_base, self.directory_dataset) |
| 435 | ) |
| 436 | |
| 437 | for ds in self.size_buckets: |
| 438 | ds.cache_latents(map_fn, regenerate_cache=regenerate_cache, trust_cache=trust_cache, caching_batch_size=caching_batch_size) |
| 439 | |
| 440 | def cache_text_embeddings(self, map_fn, i, regenerate_cache=False, caching_batch_size=1): |
| 441 | print(f'caching text embeddings: {self.ar_frames}') |
| 442 | te_dataset = _cache_text_embeddings(self.metadata_dataset, map_fn, i, self.cache_dir, regenerate_cache, caching_batch_size) |
| 443 | for size_bucket_dataset in self.size_buckets: |
| 444 | size_bucket_dataset.add_text_embedding_dataset(te_dataset) |
| 445 | |
| 446 | |
| 447 | class DirectoryDataset: |