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

Class ARBucketDataset

utils/dataset.py:398–444  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

396
397
398class 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
447class DirectoryDataset:

Callers 1

cache_metadataMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected