MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / SizeBucketDataset

Class SizeBucketDataset

utils/dataset.py:206–336  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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.
206class 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:

Callers 2

cache_latentsMethod · 0.85
cache_metadataMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected