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

Method __init__

utils/dataset.py:207–231  ·  view source on GitHub ↗
(self, metadata_dataset, directory_config, size_bucket, cache_base, directory_dataset)

Source from the content-addressed store, hash-verified

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}')

Callers

nothing calls this directly

Calls 3

seed_from_hashFunction · 0.85
bucket_suffixFunction · 0.85
getMethod · 0.80

Tested by

no test coverage detected