Skip files where caching fails in a dataset instead of crashing.
This commit is contained in:
@@ -1968,11 +1968,13 @@ class LatentCachingMixin:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
print_acc(f"Error processing image: {prep_item.path}")
|
print_acc(f"Error processing image: {prep_item.path}")
|
||||||
print_acc(f"Error: {str(e)}")
|
print_acc(f"Error: {str(e)}")
|
||||||
raise e
|
print_acc(" - Skipping file and removing it from the dataset")
|
||||||
|
return prep_item, None, None, False
|
||||||
return prep_item, prep_latent_path, None, True
|
return prep_item, prep_latent_path, None, True
|
||||||
|
|
||||||
# use tqdm to show progress
|
# use tqdm to show progress
|
||||||
i = 0
|
i = 0
|
||||||
|
failed_items = []
|
||||||
pbar = tqdm(total=len(self.file_list), desc=f'Caching latents{" to disk" if to_disk else ""}')
|
pbar = tqdm(total=len(self.file_list), desc=f'Caching latents{" to disk" if to_disk else ""}')
|
||||||
executor = ThreadPoolExecutor(max_workers=num_workers)
|
executor = ThreadPoolExecutor(max_workers=num_workers)
|
||||||
try:
|
try:
|
||||||
@@ -1986,6 +1988,11 @@ class LatentCachingMixin:
|
|||||||
next_item = next(file_iter, None)
|
next_item = next(file_iter, None)
|
||||||
if next_item is not None:
|
if next_item is not None:
|
||||||
pending.append(executor.submit(_prep, next_item))
|
pending.append(executor.submit(_prep, next_item))
|
||||||
|
if latent_path is None:
|
||||||
|
# file failed to load; drop it from the dataset and keep going
|
||||||
|
failed_items.append(file_item)
|
||||||
|
pbar.update(1)
|
||||||
|
continue
|
||||||
if needs_encode and not did_move:
|
if needs_encode and not did_move:
|
||||||
self.sd.set_device_state_preset('cache_latents')
|
self.sd.set_device_state_preset('cache_latents')
|
||||||
did_move = True
|
did_move = True
|
||||||
@@ -1997,10 +2004,33 @@ class LatentCachingMixin:
|
|||||||
executor.shutdown(wait=True, cancel_futures=True)
|
executor.shutdown(wait=True, cancel_futures=True)
|
||||||
pbar.close()
|
pbar.close()
|
||||||
|
|
||||||
|
if failed_items:
|
||||||
|
print_acc(f"Removed {len(failed_items)} files from the dataset that failed to load")
|
||||||
|
self._remove_file_items(failed_items)
|
||||||
|
|
||||||
# restore device state
|
# restore device state
|
||||||
if did_move:
|
if did_move:
|
||||||
self.sd.restore_device_state()
|
self.sd.restore_device_state()
|
||||||
|
|
||||||
|
def _remove_file_items(self: 'AiToolkitDataset', items_to_remove: List['FileItemDTO']):
|
||||||
|
# buckets hold raw indices into file_list, so removal requires remapping them
|
||||||
|
remove_ids = {id(item) for item in items_to_remove}
|
||||||
|
old_to_new = {}
|
||||||
|
new_file_list = []
|
||||||
|
for old_idx, item in enumerate(self.file_list):
|
||||||
|
if id(item) in remove_ids:
|
||||||
|
continue
|
||||||
|
old_to_new[old_idx] = len(new_file_list)
|
||||||
|
new_file_list.append(item)
|
||||||
|
self.file_list = new_file_list
|
||||||
|
if self.dataset_config.buckets and getattr(self, 'buckets', None):
|
||||||
|
for key in list(self.buckets.keys()):
|
||||||
|
bucket = self.buckets[key]
|
||||||
|
bucket.file_list_idx = [old_to_new[idx] for idx in bucket.file_list_idx if idx in old_to_new]
|
||||||
|
if len(bucket.file_list_idx) == 0:
|
||||||
|
del self.buckets[key]
|
||||||
|
self.build_batch_indices()
|
||||||
|
|
||||||
def _cache_one_latent(
|
def _cache_one_latent(
|
||||||
self: 'AiToolkitDataset',
|
self: 'AiToolkitDataset',
|
||||||
file_item: 'FileItemDTO',
|
file_item: 'FileItemDTO',
|
||||||
|
|||||||
Reference in New Issue
Block a user