| 204 | # The smallest unit of a dataset. Represents a single size bucket from a single folder of images |
| 205 | # and captions on disk. Not batched; returns individual items. |
| 206 | class SizeBucketDataset: |
| 207 | def __init__(self, metadata_dataset, directory_config, size_bucket, cache_base, directory_dataset): |
| 208 | # Shuffle deterministically based on size bucket, so that two resolutions of the same aspect ratio get different |
| 209 | # orders, which mixes data better when training on multiple resolutions at once. |
| 210 | seed = seed_from_hash(size_bucket) |
| 211 | metadata_dataset = metadata_dataset.shuffle(seed=seed) |
| 212 | self.metadata_dataset = metadata_dataset |
| 213 | self.directory_config = directory_config |
| 214 | self.size_bucket = size_bucket |
| 215 | self.path = Path(self.directory_config['path']) |
| 216 | self.cache_dir = cache_base / f'cache_{bucket_suffix(size_bucket)}' |
| 217 | self.captions_dict = directory_dataset.captions_dict # optional |
| 218 | |
| 219 | if len(size_bucket) == 4: |
| 220 | # rename old folder name to the new one for convenience |
| 221 | old_cache_dir = cache_base / f'cache_{bucket_suffix(size_bucket[1:])}' |
| 222 | if old_cache_dir.exists() and not self.cache_dir.exists(): |
| 223 | old_cache_dir.rename(self.cache_dir) |
| 224 | |
| 225 | os.makedirs(self.cache_dir, exist_ok=True) |
| 226 | self.text_embedding_datasets = [] |
| 227 | self.uncond_text_embeddings = [] |
| 228 | self.num_repeats = self.directory_config['num_repeats'] |
| 229 | self.shuffle_skip = max(directory_config.get('cache_shuffle_num', 0), 1) # Should be provided in DirectoryDataset |
| 230 | if self.num_repeats <= 0: |
| 231 | raise ValueError(f'num_repeats must be >0, was {self.num_repeats}') |
| 232 | |
| 233 | def cache_latents(self, map_fn, regenerate_cache=False, trust_cache=False, caching_batch_size=1): |
| 234 | print(f'caching latents: {self.size_bucket}') |
| 235 | self.latent_dataset = _map_and_cache( |
| 236 | self.metadata_dataset, |
| 237 | map_fn, |
| 238 | self.cache_dir, |
| 239 | cache_file_prefix='latents_', |
| 240 | regenerate_cache=regenerate_cache, |
| 241 | caching_batch_size=caching_batch_size, |
| 242 | ) |
| 243 | assert len(self.latent_dataset) == len(self.metadata_dataset), (len(self.latent_dataset), len(self.metadata_dataset)) |
| 244 | |
| 245 | iteration_order_cache_dir = self.cache_dir / 'iteration_order' |
| 246 | |
| 247 | if regenerate_cache or not iteration_order_cache_dir.exists() or not trust_cache: |
| 248 | print('Building iteration order') |
| 249 | image_spec_to_latents_idx = { |
| 250 | tuple(image_spec): i |
| 251 | for i, image_spec in enumerate(self.metadata_dataset['image_spec']) |
| 252 | } |
| 253 | |
| 254 | equal_num_captions = True |
| 255 | num_captions = None |
| 256 | for example in self.metadata_dataset.select_columns(['caption']): |
| 257 | n = len(example['caption']) |
| 258 | if num_captions is not None and n != num_captions: |
| 259 | equal_num_captions = False |
| 260 | break |
| 261 | num_captions = n |
| 262 | |
| 263 | if equal_num_captions: |
no outgoing calls
no test coverage detected