| |
| |
| import math |
| import numpy as np |
| import torch |
| from diffusers.utils import logging |
| from einops import rearrange |
|
|
| from .normalize import get_normalize_transform, get_denormalize_transform |
|
|
| logger = logging.get_logger(__name__) |
|
|
|
|
| class VAEProcessor: |
|
|
| def __init__( |
| self, |
| *, |
| vae_ratio, |
| vae_ratio_t, |
| clip_length, |
| frame_overlap, |
| token_overlap, |
| tokens_chunk_size, |
| isolated_last_frame, |
| latent_patch_size, |
| crop_mode, |
| pixel_norm_type="imagenet", |
| transform=None, |
| transform_rev=None, |
| use_3d_conv=False, |
| ): |
| self.vae_ratio = vae_ratio |
| self.vae_ratio_t = vae_ratio_t |
| self.clip_length = clip_length |
| self.frame_overlap = frame_overlap |
| self.token_overlap = token_overlap |
| self.tokens_chunk_size = tokens_chunk_size |
| self.isolated_last_frame = isolated_last_frame |
| self.latent_patch_size = latent_patch_size |
| self.crop_mode = crop_mode |
| self.transform = transform or get_normalize_transform(pixel_norm_type) |
| self.transform_rev = transform_rev or get_denormalize_transform(pixel_norm_type) |
| self.use_3d_conv = use_3d_conv |
|
|
| def _ensure_list(self, data): |
| return data if isinstance(data, list) else [data] |
|
|
| def _align_to_total_patch_size(self, h, w): |
| total_patch_size = self.latent_patch_size * self.vae_ratio |
| new_h = (h // total_patch_size) * total_patch_size |
| new_w = (w // total_patch_size) * total_patch_size |
| return new_h, new_w |
|
|
| def _crop_to_align(self, tensor, new_h, new_w, is_video=False): |
| if is_video: |
| _, _, _, h, w = tensor.shape |
| else: |
| _, _, h, w = tensor.shape |
|
|
| if self.crop_mode == "center": |
| top = (h - new_h) // 2 |
| left = (w - new_w) // 2 |
| else: |
| top = 0 |
| left = 0 |
|
|
| if is_video: |
| return tensor[:, :, :, top : top + new_h, left : left + new_w] |
| else: |
| return tensor[:, :, top : top + new_h, left : left + new_w] |
|
|
| def _align_target_token(self, T, mode): |
| intra_tail = self.clip_length % self.vae_ratio_t |
| min_frames = intra_tail or self.vae_ratio_t |
| full_chunks = T // self.clip_length |
| remainder = T % self.clip_length |
|
|
| if remainder == 0: |
| return max(T, min_frames) |
|
|
| if mode == "pad": |
| aligned_r = ( |
| math.ceil((remainder - intra_tail) / self.vae_ratio_t) * self.vae_ratio_t |
| + intra_tail |
| ) |
| if aligned_r > self.clip_length: |
| return (full_chunks + 1) * self.clip_length + intra_tail |
| return full_chunks * self.clip_length + aligned_r |
| else: |
| k = (remainder - intra_tail) // self.vae_ratio_t |
| if k >= 0: |
| target = full_chunks * self.clip_length + k * self.vae_ratio_t + intra_tail |
| return max(target, min_frames) |
| elif full_chunks > 0: |
| return full_chunks * self.clip_length |
| else: |
| return min_frames |
|
|
| def _align_target(self, T, mode, granularity): |
| if granularity == "chunk": |
| step = self.clip_length |
| tail = self.frame_overlap |
| if self.isolated_last_frame: |
| tail += 1 |
|
|
| k = math.ceil((T - tail) / step) if mode == "pad" else (T - tail) // step |
| return max(k, 1) * step + tail |
|
|
| isolated_extra = 1 if self.isolated_last_frame else 0 |
| return self._align_target_token(T - isolated_extra, mode) + isolated_extra |
|
|
| def align_video_length(self, video_length, mode="pad", granularity="chunk"): |
| target = self._align_target(video_length, mode, granularity) |
| delta = target - video_length |
| if delta > 0 and mode == "trim": |
| raise ValueError( |
| f"Cannot trim {video_length} frames to valid length {target}: " |
| f"not enough frames (granularity={granularity})" |
| ) |
| return delta |
|
|
| def align_video_length_2pass(self, video_length): |
| """Return the leading/trailing frame pads and trailing latent drop. |
| |
| This is the continuation-prefix (2-pass) alignment. The caller temporarily disables the model's normal token |
| drop and keeps these mirrored processor fields at zero. |
| """ |
| if self.isolated_last_frame: |
| raise ValueError( |
| "align_video_length_2pass does not support isolated_last_frame" |
| ) |
| if self.token_overlap != 0 or self.frame_overlap != 0: |
| raise ValueError( |
| "align_video_length_2pass requires token_drop=0 alignment" |
| ) |
|
|
| leading = self.align_video_length( |
| video_length, mode="pad", granularity="token" |
| ) |
| token_aligned = video_length + leading |
| trailing = self.align_video_length( |
| token_aligned, mode="pad", granularity="chunk" |
| ) |
|
|
| if trailing > 0: |
| intra_tail = self.clip_length % self.vae_ratio_t |
| full_chunks = token_aligned // self.clip_length |
| remainder = token_aligned % self.clip_length |
| real_tokens = full_chunks * self.tokens_chunk_size |
| if remainder > 0: |
| real_tokens += ( |
| (remainder - intra_tail) // self.vae_ratio_t + 1 |
| ) |
| drop_tokens = ( |
| self.get_latent_length(token_aligned + trailing) - real_tokens |
| ) |
| else: |
| drop_tokens = 0 |
|
|
| return leading, trailing, drop_tokens |
|
|
| def get_suitable_video_length(self, video_length, verbose=False): |
| used_frame_length = video_length + self.align_video_length( |
| video_length, mode="trim", granularity="chunk" |
| ) |
| if verbose: |
| logger.info( |
| f"Pick first {used_frame_length} frames from {video_length}-frame video" |
| ) |
| return used_frame_length |
|
|
| def get_latent_length(self, video_length): |
| tail_frame = self.frame_overlap |
| tail_token = self.token_overlap |
| if self.isolated_last_frame: |
| tail_frame += 1 |
| tail_token += 1 |
|
|
| video_length = self.get_suitable_video_length(video_length) |
| latent_length = ( |
| int((video_length - tail_frame) // self.clip_length) |
| * self.tokens_chunk_size |
| + tail_token |
| ) |
| return latent_length |
|
|
|
|
|
|
| def transform_tensor(self, tensor): |
| B, T = None, None |
| if tensor.ndim == 5: |
| if tensor.shape[2] == 3: |
| tensor = tensor.transpose(1, 2) |
| B, _, T, _, _ = tensor.shape |
| tensor = rearrange(tensor, "b c t h w -> (b t) c h w") |
| elif tensor.ndim == 4: |
| if tensor.shape[0] == 3: |
| tensor = tensor.transpose(0, 1) |
| elif tensor.ndim == 3: |
| tensor = tensor.unsqueeze(0) |
| else: |
| raise ValueError(f"Unsupported tensor shape: {tensor.shape}") |
|
|
| tensor = self.transform(tensor) |
|
|
| if B is not None and T is not None: |
| tensor = rearrange(tensor, "(b t) c h w -> b c t h w", b=B, t=T) |
|
|
| return tensor.contiguous() |
|
|
| def revert_tensor(self, tensor): |
| B, T = None, None |
| if self.use_3d_conv: |
| tensor = tensor.unsqueeze(2) if tensor.ndim == 4 else tensor |
| B, _, T, _, _ = tensor.shape |
| tensor = rearrange(tensor, "b c t h w -> (b t) c h w") |
| tensor_rev = self.transform_rev(tensor).clamp(0, 1) |
| if B is not None: |
| tensor_rev = rearrange(tensor_rev, "(b t) c h w -> b c t h w", b=B, t=T) |
| return tensor_rev.contiguous() |
|
|
| @staticmethod |
| def convert_numpy_to_tensor(numpy_array, device=None): |
| if isinstance(numpy_array, list): |
| numpy_array = np.stack(numpy_array, axis=0) |
| numpy_array = numpy_array.astype(np.float32) |
| tensor = torch.from_numpy(numpy_array) |
| tensor = tensor.permute(0, 3, 1, 2) |
| tensor = tensor / 255.0 |
| if device is not None: |
| tensor = tensor.to(device) |
| return tensor |
|
|
|
|