Improvements to video frame loading. Added ability to cache as uint8 pixelspace for video

This commit is contained in:
Jaret Burkett
2026-07-24 12:36:47 -06:00
parent 67984754c3
commit 8f2d001eae

View File

@@ -552,24 +552,34 @@ class ImageProcessingDTOMixin:
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
print_acc(f" Frames to extract: {frames_to_extract}")
# Extract frames
# Extract frames -- decode sequentially in a single pass. A cap.set() seek per
# frame forces a keyframe seek + GOP re-decode for every extracted frame
# (~20x slower); frames_to_extract is always ascending, so seek once to the
# first frame then grab() through the gaps.
frames = []
for frame_idx in frames_to_extract:
# Safety check - ensure frame_idx is within bounds (silently fix)
if frame_idx > max_frame_index:
frame_idx = max_frame_index
# Set frame position
cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)
unique_frame_idxs = sorted(set(frames_to_extract))
processed_frames = {} # frame_idx -> processed frame (duplicates reuse it)
# Set frame position
pos = unique_frame_idxs[0]
if pos > 0:
cap.set(cv2.CAP_PROP_POS_FRAMES, pos)
# Silently verify position was set correctly (no warnings unless debug mode)
if hasattr(self.dataset_config, 'debug') and self.dataset_config.debug:
actual_pos = int(cap.get(cv2.CAP_PROP_POS_FRAMES))
if actual_pos != frame_idx:
print_acc(f"Warning: Failed to set exact frame position. Requested: {frame_idx}, Actual: {actual_pos}")
if actual_pos != pos:
print_acc(f"Warning: Failed to set exact frame position. Requested: {pos}, Actual: {actual_pos}")
for frame_idx in unique_frame_idxs:
# skip past frames between targets without decoding them to images
while pos < frame_idx and cap.grab():
pos += 1
ret, frame = cap.read()
if not ret:
if ret:
pos += 1
else:
# Try to provide more detailed error information
actual_frame = int(cap.get(cv2.CAP_PROP_POS_FRAMES))
frame_pos_info = f"Requested frame: {frame_idx}, Actual frame position: {actual_frame}"
@@ -586,6 +596,8 @@ class ImageProcessingDTOMixin:
print_acc(f"Falling back to nearby frame {fallback_pos} instead of {frame_idx}")
frame = fallback_frame
fallback_success = True
# resync sequential position after the fallback seek
pos = fallback_pos + 1
break
else:
# No fallback worked, raise a more detailed exception
@@ -618,9 +630,12 @@ class ImageProcessingDTOMixin:
# Apply transform if provided
if transform:
img = transform(img)
frames.append(img)
processed_frames[frame_idx] = img
# assemble in extraction order; stretched clips repeat decoded frames
frames = [processed_frames[frame_idx] for frame_idx in frames_to_extract]
# Release the video capture
cap.release()
@@ -1609,6 +1624,16 @@ class ArgBreakMixin:
pass
def _latent_to_uint8(latent: torch.Tensor) -> torch.Tensor:
# pixel-space latents in [-1, 1] -> uint8 0..255 for compact caching
return ((latent.float().clamp(-1, 1) + 1.0) * 127.5).round().to(torch.uint8)
def _latent_from_uint8(latent: torch.Tensor, dtype: torch.dtype = torch.float32) -> torch.Tensor:
# uint8 0..255 -> pixel-space latents in [-1, 1]
return (latent.to(torch.float32) / 127.5 - 1.0).to(dtype)
class LatentCachingFileItemDTOMixin:
def __init__(self, *args, **kwargs):
# if we have super, call it
@@ -1709,8 +1734,13 @@ class LatentCachingFileItemDTOMixin:
device='cpu'
)
self._encoded_latent = state_dict['latent']
if self._encoded_latent.dtype == torch.uint8:
# pixel-space latents cached as uint8
self._encoded_latent = _latent_from_uint8(self._encoded_latent)
if 'first_frame_latent' in state_dict:
self._cached_first_frame_latent = state_dict['first_frame_latent']
if self._cached_first_frame_latent.dtype == torch.uint8:
self._cached_first_frame_latent = _latent_from_uint8(self._cached_first_frame_latent)
if 'audio_latent' in state_dict:
self._cached_audio_latent = state_dict['audio_latent']
if 'num_frames' in state_dict:
@@ -1752,9 +1782,16 @@ class LatentCachingMixin:
if to_memory:
# load it into memory
state_dict = load_file(latent_path, device='cpu')
file_item._encoded_latent = state_dict['latent'].to('cpu', dtype=self.sd.torch_dtype)
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:
file_item._cached_first_frame_latent = state_dict['first_frame_latent'].to('cpu', dtype=self.sd.torch_dtype)
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:
@@ -1768,11 +1805,15 @@ class LatentCachingMixin:
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:
state_dict['latent'] = latent.clone().detach().cpu()
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)}")
@@ -1789,7 +1830,10 @@ class LatentCachingMixin:
raise ValueError(f"Unknown frame shape {frames.shape}")
first_frame_latent = self.sd.encode_images(first_frames).squeeze(0)
if to_disk:
state_dict['first_frame_latent'] = first_frame_latent.clone().detach().cpu()
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: