Queue up videos with multiple threads when caching latents so the VAE is not waiting on videos to process
This commit is contained in:
@@ -1030,6 +1030,8 @@ class DatasetConfig:
|
||||
|
||||
self.num_workers: int = kwargs.get('num_workers', 2)
|
||||
self.prefetch_factor: int = kwargs.get('prefetch_factor', 2)
|
||||
# threads used to prep (decode/resize) items ahead of the VAE while caching latents
|
||||
self.cache_latents_num_workers: int = kwargs.get('cache_latents_num_workers', min(6, os.cpu_count() or 1))
|
||||
self.extra_values: List[float] = kwargs.get('extra_values', [])
|
||||
self.square_crop: bool = kwargs.get('square_crop', False)
|
||||
# apply same augmentations to control images. Usually want this true unless special case
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
import base64
|
||||
import glob
|
||||
import hashlib
|
||||
import itertools
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import OrderedDict
|
||||
from collections import OrderedDict, deque
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import TYPE_CHECKING, List, Dict, Union
|
||||
import traceback
|
||||
|
||||
@@ -1812,111 +1814,154 @@ class LatentCachingMixin:
|
||||
# move sd items to cpu except for vae
|
||||
self.sd.set_device_state_preset('cache_latents')
|
||||
|
||||
# prep (video decode, frame extraction, audio load, disk reads) is done by a
|
||||
# thread pool so the next items are ready while the current one is encoding.
|
||||
# the in-flight window is bounded so decoded videos don't pile up in RAM.
|
||||
num_workers = max(1, self.dataset_config.cache_latents_num_workers)
|
||||
|
||||
def _prep(prep_item: 'FileItemDTO'):
|
||||
prep_item.is_caching_to_disk = to_disk
|
||||
prep_item.is_caching_to_memory = to_memory
|
||||
prep_item.latent_load_device = self.sd.device
|
||||
|
||||
prep_latent_path = prep_item.get_latent_path(recalculate=True)
|
||||
try:
|
||||
if os.path.exists(prep_latent_path):
|
||||
cached_state_dict = load_file(prep_latent_path, device='cpu') if to_memory else None
|
||||
return prep_item, prep_latent_path, cached_state_dict, False
|
||||
# not saved to disk, load the image/video/audio
|
||||
prep_item.load_and_process_image(self.transform, only_load_latents=True)
|
||||
except Exception as e:
|
||||
print_acc(f"Error processing image: {prep_item.path}")
|
||||
print_acc(f"Error: {str(e)}")
|
||||
raise e
|
||||
return prep_item, prep_latent_path, None, True
|
||||
|
||||
# use tqdm to show progress
|
||||
i = 0
|
||||
for file_item in tqdm(self.file_list, desc=f'Caching latents{" to disk" if to_disk else ""}'):
|
||||
file_item.is_caching_to_disk = to_disk
|
||||
file_item.is_caching_to_memory = to_memory
|
||||
file_item.latent_load_device = self.sd.device
|
||||
|
||||
latent_path = file_item.get_latent_path(recalculate=True)
|
||||
# check if it is saved to disk already
|
||||
if os.path.exists(latent_path):
|
||||
if to_memory:
|
||||
# load it into memory
|
||||
state_dict = load_file(latent_path, device='cpu')
|
||||
cached_latent = state_dict['latent']
|
||||
if cached_latent.dtype == torch.uint8:
|
||||
# pixel-space latents cached as uint8
|
||||
cached_latent = _latent_from_uint8(cached_latent)
|
||||
file_item._encoded_latent = cached_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if 'first_frame_latent' in state_dict:
|
||||
cached_first_frame = state_dict['first_frame_latent']
|
||||
if cached_first_frame.dtype == torch.uint8:
|
||||
cached_first_frame = _latent_from_uint8(cached_first_frame)
|
||||
file_item._cached_first_frame_latent = cached_first_frame.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if 'audio_latent' in state_dict:
|
||||
file_item._cached_audio_latent = state_dict['audio_latent'].to('cpu', dtype=self.sd.torch_dtype)
|
||||
else:
|
||||
# not saved to disk, calculate
|
||||
# load the image first
|
||||
file_item.load_and_process_image(self.transform, only_load_latents=True)
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
state_dict = OrderedDict()
|
||||
first_frame_latent = None
|
||||
audio_latent = None
|
||||
frames = None
|
||||
# add batch dimension
|
||||
cache_uint8 = getattr(self.sd, 'cache_latents_as_uint8', False)
|
||||
try:
|
||||
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
latent = self.sd.encode_images(imgs).squeeze(0)
|
||||
if to_disk:
|
||||
if cache_uint8:
|
||||
state_dict['latent'] = _latent_to_uint8(latent).cpu()
|
||||
else:
|
||||
state_dict['latent'] = latent.clone().detach().cpu()
|
||||
except Exception as e:
|
||||
print_acc(f"Error processing image: {file_item.path}")
|
||||
print_acc(f"Error: {str(e)}")
|
||||
raise e
|
||||
# do first frame
|
||||
is_video = self.dataset_config.auto_frame_count or self.dataset_config.num_frames > 1
|
||||
if is_video and self.dataset_config.do_i2v:
|
||||
frames = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
if len(frames.shape) == 4:
|
||||
first_frames = frames
|
||||
elif len(frames.shape) == 5:
|
||||
first_frames = frames[:, 0]
|
||||
else:
|
||||
raise ValueError(f"Unknown frame shape {frames.shape}")
|
||||
first_frame_latent = self.sd.encode_images(first_frames).squeeze(0)
|
||||
if to_disk:
|
||||
if cache_uint8:
|
||||
state_dict['first_frame_latent'] = _latent_to_uint8(first_frame_latent).cpu()
|
||||
else:
|
||||
state_dict['first_frame_latent'] = first_frame_latent.clone().detach().cpu()
|
||||
|
||||
# audio (video+audio models only — audio-only models already encoded above via encode_images)
|
||||
if not self.is_audio_model and file_item.audio_data is not None:
|
||||
audio_latent = self.sd.encode_audio([file_item.audio_data]).squeeze(0)
|
||||
if to_disk:
|
||||
state_dict['audio_latent'] = audio_latent.clone().detach().cpu()
|
||||
|
||||
if is_video:
|
||||
state_dict['num_frames'] = torch.tensor(file_item.num_frames, dtype=torch.int32)
|
||||
|
||||
# save_latent
|
||||
if to_disk:
|
||||
# metadata
|
||||
meta = get_meta_for_safetensors(file_item.get_latent_info_dict())
|
||||
os.makedirs(os.path.dirname(latent_path), exist_ok=True)
|
||||
save_file(state_dict, latent_path, metadata=meta)
|
||||
|
||||
if to_memory:
|
||||
# keep it in memory
|
||||
file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if first_frame_latent is not None:
|
||||
file_item._cached_first_frame_latent = first_frame_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if audio_latent is not None:
|
||||
file_item._cached_audio_latent = audio_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
|
||||
del imgs
|
||||
del latent
|
||||
del frames
|
||||
del file_item.tensor
|
||||
del state_dict
|
||||
del first_frame_latent
|
||||
del audio_latent
|
||||
file_item.cleanup()
|
||||
|
||||
file_item.is_latent_cached = True
|
||||
i += 1
|
||||
pbar = tqdm(total=len(self.file_list), desc=f'Caching latents{" to disk" if to_disk else ""}')
|
||||
executor = ThreadPoolExecutor(max_workers=num_workers)
|
||||
try:
|
||||
pending = deque()
|
||||
file_iter = iter(self.file_list)
|
||||
for queued_item in itertools.islice(file_iter, num_workers + 2):
|
||||
pending.append(executor.submit(_prep, queued_item))
|
||||
while pending:
|
||||
file_item, latent_path, cached_state_dict, needs_encode = pending.popleft().result()
|
||||
# keep the window full
|
||||
next_item = next(file_iter, None)
|
||||
if next_item is not None:
|
||||
pending.append(executor.submit(_prep, next_item))
|
||||
self._cache_one_latent(file_item, latent_path, cached_state_dict, needs_encode, to_disk, to_memory)
|
||||
file_item.is_latent_cached = True
|
||||
i += 1
|
||||
pbar.update(1)
|
||||
finally:
|
||||
executor.shutdown(wait=True, cancel_futures=True)
|
||||
pbar.close()
|
||||
|
||||
# restore device state
|
||||
self.sd.restore_device_state()
|
||||
|
||||
def _cache_one_latent(
|
||||
self: 'AiToolkitDataset',
|
||||
file_item: 'FileItemDTO',
|
||||
latent_path: str,
|
||||
cached_state_dict,
|
||||
needs_encode: bool,
|
||||
to_disk: bool,
|
||||
to_memory: bool,
|
||||
):
|
||||
# check if it is saved to disk already
|
||||
if not needs_encode:
|
||||
if to_memory:
|
||||
# load it into memory
|
||||
state_dict = cached_state_dict
|
||||
cached_latent = state_dict['latent']
|
||||
if cached_latent.dtype == torch.uint8:
|
||||
# pixel-space latents cached as uint8
|
||||
cached_latent = _latent_from_uint8(cached_latent)
|
||||
file_item._encoded_latent = cached_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if 'first_frame_latent' in state_dict:
|
||||
cached_first_frame = state_dict['first_frame_latent']
|
||||
if cached_first_frame.dtype == torch.uint8:
|
||||
cached_first_frame = _latent_from_uint8(cached_first_frame)
|
||||
file_item._cached_first_frame_latent = cached_first_frame.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if 'audio_latent' in state_dict:
|
||||
file_item._cached_audio_latent = state_dict['audio_latent'].to('cpu', dtype=self.sd.torch_dtype)
|
||||
else:
|
||||
# not saved to disk, calculate
|
||||
# the image/video/audio was already loaded by the prep thread
|
||||
dtype = self.sd.torch_dtype
|
||||
device = self.sd.device_torch
|
||||
state_dict = OrderedDict()
|
||||
first_frame_latent = None
|
||||
audio_latent = None
|
||||
frames = None
|
||||
# add batch dimension
|
||||
cache_uint8 = getattr(self.sd, 'cache_latents_as_uint8', False)
|
||||
try:
|
||||
imgs = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
latent = self.sd.encode_images(imgs).squeeze(0)
|
||||
if to_disk:
|
||||
if cache_uint8:
|
||||
state_dict['latent'] = _latent_to_uint8(latent).cpu()
|
||||
else:
|
||||
state_dict['latent'] = latent.clone().detach().cpu()
|
||||
except Exception as e:
|
||||
print_acc(f"Error processing image: {file_item.path}")
|
||||
print_acc(f"Error: {str(e)}")
|
||||
raise e
|
||||
# do first frame
|
||||
is_video = self.dataset_config.auto_frame_count or self.dataset_config.num_frames > 1
|
||||
if is_video and self.dataset_config.do_i2v:
|
||||
frames = file_item.tensor.unsqueeze(0).to(device, dtype=dtype)
|
||||
if len(frames.shape) == 4:
|
||||
first_frames = frames
|
||||
elif len(frames.shape) == 5:
|
||||
first_frames = frames[:, 0]
|
||||
else:
|
||||
raise ValueError(f"Unknown frame shape {frames.shape}")
|
||||
first_frame_latent = self.sd.encode_images(first_frames).squeeze(0)
|
||||
if to_disk:
|
||||
if cache_uint8:
|
||||
state_dict['first_frame_latent'] = _latent_to_uint8(first_frame_latent).cpu()
|
||||
else:
|
||||
state_dict['first_frame_latent'] = first_frame_latent.clone().detach().cpu()
|
||||
|
||||
# audio (video+audio models only — audio-only models already encoded above via encode_images)
|
||||
if not self.is_audio_model and file_item.audio_data is not None:
|
||||
audio_latent = self.sd.encode_audio([file_item.audio_data]).squeeze(0)
|
||||
if to_disk:
|
||||
state_dict['audio_latent'] = audio_latent.clone().detach().cpu()
|
||||
|
||||
if is_video:
|
||||
state_dict['num_frames'] = torch.tensor(file_item.num_frames, dtype=torch.int32)
|
||||
|
||||
# save_latent
|
||||
if to_disk:
|
||||
# metadata
|
||||
meta = get_meta_for_safetensors(file_item.get_latent_info_dict())
|
||||
os.makedirs(os.path.dirname(latent_path), exist_ok=True)
|
||||
save_file(state_dict, latent_path, metadata=meta)
|
||||
|
||||
if to_memory:
|
||||
# keep it in memory
|
||||
file_item._encoded_latent = latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if first_frame_latent is not None:
|
||||
file_item._cached_first_frame_latent = first_frame_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
if audio_latent is not None:
|
||||
file_item._cached_audio_latent = audio_latent.to('cpu', dtype=self.sd.torch_dtype)
|
||||
|
||||
del imgs
|
||||
del latent
|
||||
del frames
|
||||
del file_item.tensor
|
||||
del state_dict
|
||||
del first_frame_latent
|
||||
del audio_latent
|
||||
file_item.cleanup()
|
||||
|
||||
|
||||
class TextEmbeddingFileItemDTOMixin:
|
||||
def __init__(self, *args, **kwargs):
|
||||
|
||||
Reference in New Issue
Block a user