(self, map_fn, regenerate_cache=False, trust_cache=False, caching_batch_size=1)
| 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}') |
nothing calls this directly
no test coverage detected