diff --git a/config/lingbot_world/lingbot_world_i2v.yaml b/config/lingbot_world/lingbot_world_i2v.yaml new file mode 100644 index 0000000..4c20ed0 --- /dev/null +++ b/config/lingbot_world/lingbot_world_i2v.yaml @@ -0,0 +1,56 @@ +format: civitai +pipeline: LingBotWorld +transformer_additional_kwargs: + transformer_low_noise_model_subpath: ./low_noise_model + transformer_high_noise_model_subpath: ./high_noise_model + transformer_combination_type: "moe" + # LingBot-World uses higher boundary than Wan2.2 (0.947 vs 0.900) + boundary: 0.947 + model_type: i2v + cross_attn_type: cross_attn + dict_mapping: + in_dim: in_channels + dim: hidden_size + +vae_kwargs: + vae_type: "AutoencoderKLWan" + vae_subpath: Wan2.1_VAE.pth + temporal_compression_ratio: 4 + spatial_compression_ratio: 8 + +text_encoder_kwargs: + text_encoder_subpath: models_t5_umt5-xxl-enc-bf16.pth + tokenizer_subpath: google/umt5-xxl + text_length: 512 + vocab: 256384 + dim: 4096 + dim_attn: 4096 + dim_ffn: 10240 + num_heads: 64 + num_layers: 24 + num_buckets: 32 + shared_pos: False + dropout: 0.0 + +scheduler_kwargs: + scheduler_subpath: null + num_train_timesteps: 1000 + # LingBot-World uses higher shift than Wan2.2 (10.0 vs 5.0) + shift: 10.0 + use_dynamic_shifting: false + base_shift: 0.5 + max_shift: 1.15 + base_image_seq_len: 256 + max_image_seq_len: 4096 + +image_encoder_kwargs: + image_encoder_subpath: models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth + +# LingBot-World specific inference settings +inference_kwargs: + # Number of inference steps (LingBot-World uses 70, Wan2.2 uses 40) + num_inference_steps: 70 + # Guidance scale for both low and high noise models + guidance_scale: 5.0 + # Default negative prompt + negative_prompt: "画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,皮肤,肢体,面部特征,汽车,电线" diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 1a7a2e6..0fc6e54 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -42,6 +42,7 @@ from .wan_text_encoder import WanT5EncoderModel from .wan_transformer3d import (Wan2_2Transformer3DModel, WanRMSNorm, WanSelfAttention, WanTransformer3DModel) from .wan_transformer3d_animate import Wan2_2Transformer3DModel_Animate +from .wan_transformer3d_lingbot import LingBotWanTransformer3DModel from .wan_transformer3d_s2v import Wan2_2Transformer3DModel_S2V from .wan_transformer3d_vace import VaceWanTransformer3DModel from .wan_vae import AutoencoderKLWan, AutoencoderKLWan_ diff --git a/videox_fun/pipeline/__init__.py b/videox_fun/pipeline/__init__.py index a862a3a..32a97f9 100755 --- a/videox_fun/pipeline/__init__.py +++ b/videox_fun/pipeline/__init__.py @@ -28,6 +28,7 @@ from .pipeline_wan_phantom import WanFunPhantomPipeline from .pipeline_wan_vace import WanVacePipeline from .pipeline_z_image import ZImagePipeline from .pipeline_z_image_control import ZImageControlPipeline +from .pipeline_lingbot_world import LingBotWorldI2VPipeline WanFunPipeline = WanPipeline WanI2VPipeline = WanFunInpaintPipeline diff --git a/videox_fun/pipeline/pipeline_lingbot_world.py b/videox_fun/pipeline/pipeline_lingbot_world.py new file mode 100644 index 0000000..6c5aaf6 --- /dev/null +++ b/videox_fun/pipeline/pipeline_lingbot_world.py @@ -0,0 +1,840 @@ +import inspect +import math +import os +from dataclasses import dataclass +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +import numpy as np +import torch +import torch.nn.functional as TorchF +import torchvision.transforms.functional as TF +from diffusers import FlowMatchEulerDiscreteScheduler +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.image_processor import VaeImageProcessor +from diffusers.pipelines.pipeline_utils import DiffusionPipeline +from diffusers.utils import BaseOutput, logging, replace_example_docstring +from diffusers.utils.torch_utils import randn_tensor +from diffusers.video_processor import VideoProcessor +from einops import rearrange +from PIL import Image + +from ..models import AutoencoderKLWan, AutoTokenizer, WanT5EncoderModel +from ..models.wan_transformer3d_lingbot import LingBotWanTransformer3DModel +from ..utils.cam_utils import ( + compute_relative_poses, + get_Ks_transformed, + get_plucker_embeddings, + interpolate_camera_poses, +) +from ..utils.fm_solvers import FlowDPMSolverMultistepScheduler, get_sampling_sigmas +from ..utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + +logger = logging.get_logger(__name__) + + +EXAMPLE_DOC_STRING = """ + Examples: + ```python + # LingBot-World I2V with camera control + from videox_fun.pipeline import LingBotWorldI2VPipeline + + pipeline = LingBotWorldI2VPipeline(...) + video = pipeline( + prompt="A beautiful landscape", + video=input_video, + mask_video=mask, + action_path="path/to/camera/data", # Optional camera control + ) + ``` +""" + + +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + sigmas: Optional[List[float]] = None, + **kwargs, +): + """Retrieve timesteps from scheduler.""" + if timesteps is not None and sigmas is not None: + raise ValueError("Only one of `timesteps` or `sigmas` can be passed.") + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + elif sigmas is not None: + accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accept_sigmas: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" sigmas schedules." + ) + scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +def resize_mask(mask, latent, process_first_frame_only=True): + """Resize mask to match latent dimensions.""" + latent_size = latent.size() + batch_size, channels, num_frames, height, width = mask.shape + + if process_first_frame_only: + target_size = list(latent_size[2:]) + target_size[0] = 1 + first_frame_resized = TorchF.interpolate( + mask[:, :, 0:1, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + + target_size = list(latent_size[2:]) + target_size[0] = target_size[0] - 1 + if target_size[0] != 0: + remaining_frames_resized = TorchF.interpolate( + mask[:, :, 1:, :, :], + size=target_size, + mode='trilinear', + align_corners=False + ) + resized_mask = torch.cat([first_frame_resized, remaining_frames_resized], dim=2) + else: + resized_mask = first_frame_resized + else: + target_size = list(latent_size[2:]) + resized_mask = TorchF.interpolate( + mask, + size=target_size, + mode='trilinear', + align_corners=False + ) + return resized_mask + + +@dataclass +class LingBotWorldPipelineOutput(BaseOutput): + """ + Output class for LingBot-World pipelines. + + Args: + videos: Generated video tensor with shape (batch_size, num_frames, channels, height, width) + """ + videos: torch.Tensor + + +class LingBotWorldI2VPipeline(DiffusionPipeline): + """ + Pipeline for image-to-video generation using LingBot-World with optional camera control. + """ + + _optional_components = ["transformer_2"] + model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae" + + _callback_tensor_inputs = [ + "latents", + "prompt_embeds", + "negative_prompt_embeds", + ] + + def __init__( + self, + tokenizer: AutoTokenizer, + text_encoder: WanT5EncoderModel, + vae: AutoencoderKLWan, + transformer: LingBotWanTransformer3DModel, + transformer_2: LingBotWanTransformer3DModel = None, + scheduler: FlowMatchEulerDiscreteScheduler = None, + ): + super().__init__() + + self.register_modules( + tokenizer=tokenizer, + text_encoder=text_encoder, + vae=vae, + transformer=transformer, + transformer_2=transformer_2, + scheduler=scheduler + ) + self.video_processor = VideoProcessor(vae_scale_factor=self.vae.spatial_compression_ratio) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae.spatial_compression_ratio) + self.mask_processor = VaeImageProcessor( + vae_scale_factor=self.vae.spatial_compression_ratio, + do_normalize=False, + do_binarize=True, + do_convert_grayscale=True + ) + + def _get_t5_prompt_embeds( + self, + prompt: Union[str, List[str]] = None, + num_videos_per_prompt: int = 1, + max_sequence_length: int = 512, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ): + device = device or self._execution_device + dtype = dtype or self.text_encoder.dtype + + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) + + text_inputs = self.tokenizer( + prompt, + padding="max_length", + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + prompt_attention_mask = text_inputs.attention_mask + untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because `max_sequence_length` is set to " + f" {max_sequence_length} tokens: {removed_text}" + ) + + seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() + prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return [u[:v] for u, v in zip(prompt_embeds, seq_lens)] + + def encode_prompt( + self, + prompt: Union[str, List[str]], + negative_prompt: Optional[Union[str, List[str]]] = None, + do_classifier_free_guidance: bool = True, + num_videos_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + max_sequence_length: int = 512, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ): + """Encodes the prompt into text encoder hidden states.""" + device = device or self._execution_device + + prompt = [prompt] if isinstance(prompt, str) else prompt + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds = self._get_t5_prompt_embeds( + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + if do_classifier_free_guidance and negative_prompt_embeds is None: + negative_prompt = negative_prompt or "" + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}." + ) + + negative_prompt_embeds = self._get_t5_prompt_embeds( + prompt=negative_prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + return prompt_embeds, negative_prompt_embeds + + def prepare_latents( + self, batch_size, num_channels_latents, num_frames, height, width, dtype, device, generator, latents=None + ): + """Prepare latent tensors.""" + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}." + ) + + shape = ( + batch_size, + num_channels_latents, + (num_frames - 1) // self.vae.temporal_compression_ratio + 1, + height // self.vae.spatial_compression_ratio, + width // self.vae.spatial_compression_ratio, + ) + + if latents is None: + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + else: + latents = latents.to(device) + + if hasattr(self.scheduler, "init_noise_sigma"): + latents = latents * self.scheduler.init_noise_sigma + return latents + + def prepare_mask_latents( + self, mask, masked_image, batch_size, height, width, dtype, device, generator, do_classifier_free_guidance, noise_aug_strength + ): + """Prepare mask latent variables.""" + if mask is not None: + mask = mask.to(device=device, dtype=self.vae.dtype) + bs = 1 + new_mask = [] + for i in range(0, mask.shape[0], bs): + mask_bs = mask[i : i + bs] + mask_bs = self.vae.encode(mask_bs)[0] + mask_bs = mask_bs.mode() + new_mask.append(mask_bs) + mask = torch.cat(new_mask, dim=0) + + if masked_image is not None: + masked_image = masked_image.to(device=device, dtype=self.vae.dtype) + bs = 1 + new_mask_pixel_values = [] + for i in range(0, masked_image.shape[0], bs): + mask_pixel_values_bs = masked_image[i : i + bs] + mask_pixel_values_bs = self.vae.encode(mask_pixel_values_bs)[0] + mask_pixel_values_bs = mask_pixel_values_bs.mode() + new_mask_pixel_values.append(mask_pixel_values_bs) + masked_image_latents = torch.cat(new_mask_pixel_values, dim=0) + else: + masked_image_latents = None + + return mask, masked_image_latents + + def prepare_camera_conditioning( + self, + action_path: str, + frame_num: int, + height: int, + width: int, + lat_f: int, + lat_h: int, + lat_w: int, + device: torch.device, + dtype: torch.dtype, + ) -> Optional[Dict[str, torch.Tensor]]: + """ + Prepare camera conditioning from action_path. + + Args: + action_path: Path to directory containing poses.npy and intrinsics.npy + frame_num: Number of video frames + height: Video height + width: Video width + lat_f: Latent temporal dimension + lat_h: Latent height + lat_w: Latent width + device: Target device + dtype: Target dtype + + Returns: + Dictionary containing c2ws_plucker_emb or None if action_path is None + """ + if action_path is None: + return None + + # Load camera data + c2ws = np.load(os.path.join(action_path, "poses.npy")) # opencv coordinate + len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1 + frame_num = min(frame_num, len_c2ws) + c2ws = c2ws[:frame_num] + + Ks = torch.from_numpy(np.load(os.path.join(action_path, "intrinsics.npy"))).float() + + # Transform intrinsics for the target resolution + # The provided intrinsics are for original image size (480p) + Ks = get_Ks_transformed( + Ks, + height_org=480, + width_org=832, + height_resize=height, + width_resize=width, + height_final=height, + width_final=width + ) + Ks = Ks[0] + + # Interpolate camera poses to latent temporal resolution + len_c2ws = len(c2ws) + c2ws_infer = interpolate_camera_poses( + src_indices=np.linspace(0, len_c2ws - 1, len_c2ws), + src_rot_mat=c2ws[:, :3, :3], + src_trans_vec=c2ws[:, :3, 3], + tgt_indices=np.linspace(0, len_c2ws - 1, int((len_c2ws - 1) // 4) + 1), + ) + c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True) + Ks = Ks.repeat(len(c2ws_infer), 1) + + c2ws_infer = c2ws_infer.to(device) + Ks = Ks.to(device) + + # Generate Plucker embeddings + c2ws_plucker_emb = get_plucker_embeddings(c2ws_infer, Ks, height, width) + c2ws_plucker_emb = rearrange( + c2ws_plucker_emb, + 'f (h c1) (w c2) c -> (f h w) (c c1 c2)', + c1=int(height // lat_h), + c2=int(width // lat_w), + ) + c2ws_plucker_emb = c2ws_plucker_emb[None, ...] # [1, f*h*w, c] + c2ws_plucker_emb = rearrange( + c2ws_plucker_emb, + 'b (f h w) c -> b c f h w', + f=lat_f, h=lat_h, w=lat_w + ).to(dtype) + + dit_cond_dict = { + "c2ws_plucker_emb": c2ws_plucker_emb.chunk(1, dim=0), + } + + return dit_cond_dict + + def decode_latents(self, latents: torch.Tensor) -> torch.Tensor: + """Decode latents to video frames.""" + frames = self.vae.decode(latents.to(self.vae.dtype)).sample + frames = (frames / 2 + 0.5).clamp(0, 1) + frames = frames.cpu().float().numpy() + return frames + + def prepare_extra_step_kwargs(self, generator, eta): + """Prepare extra kwargs for the scheduler step.""" + accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys()) + extra_step_kwargs = {} + if accepts_eta: + extra_step_kwargs["eta"] = eta + + accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys()) + if accepts_generator: + extra_step_kwargs["generator"] = generator + return extra_step_kwargs + + def check_inputs( + self, + prompt, + height, + width, + negative_prompt, + callback_on_step_end_tensor_inputs, + prompt_embeds=None, + negative_prompt_embeds=None, + ): + """Check inputs for validity.""" + if height % 8 != 0 or width % 8 != 0: + raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.") + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, " + f"but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both undefined." + ) + elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): + raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") + + if prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}." + ) + + if negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}." + ) + + if prompt_embeds is not None and negative_prompt_embeds is not None: + if prompt_embeds.shape != negative_prompt_embeds.shape: + raise ValueError( + "`prompt_embeds` and `negative_prompt_embeds` must have the same shape." + ) + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def attention_kwargs(self): + return self._attention_kwargs + + @property + def interrupt(self): + return self._interrupt + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Optional[Union[str, List[str]]] = None, + negative_prompt: Optional[Union[str, List[str]]] = None, + height: int = 480, + width: int = 832, + video: Union[torch.FloatTensor] = None, + mask_video: Union[torch.FloatTensor] = None, + num_frames: int = 81, + num_inference_steps: int = 70, + timesteps: Optional[List[int]] = None, + guidance_scale: float = 5.0, + num_videos_per_prompt: int = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: str = "numpy", + return_dict: bool = False, + callback_on_step_end: Optional[ + Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks] + ] = None, + attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + boundary: float = 0.947, + comfyui_progressbar: bool = False, + shift: int = 10, + action_path: Optional[str] = None, + ) -> Union[LingBotWorldPipelineOutput, Tuple]: + """ + Generate video from image with optional camera control. + + Args: + prompt: Text prompt for generation + negative_prompt: Negative prompt for exclusion + height: Video height + width: Video width + video: Input video tensor + mask_video: Mask video tensor + num_frames: Number of frames to generate + num_inference_steps: Number of denoising steps + timesteps: Custom timesteps + guidance_scale: Classifier-free guidance scale + num_videos_per_prompt: Number of videos per prompt + eta: Eta for DDIM scheduler + generator: Random generator + latents: Pre-generated latents + prompt_embeds: Pre-computed prompt embeddings + negative_prompt_embeds: Pre-computed negative prompt embeddings + output_type: Output format ("numpy", "latent", etc.) + return_dict: Whether to return a dict + callback_on_step_end: Callback function + attention_kwargs: Additional attention kwargs + callback_on_step_end_tensor_inputs: Tensor inputs for callback + max_sequence_length: Max text sequence length + boundary: Timestep boundary for model switching + comfyui_progressbar: Enable ComfyUI progress bar + shift: Noise schedule shift parameter + action_path: Path to camera data (poses.npy, intrinsics.npy) + + Examples: + + Returns: + Generated video tensor + """ + if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): + callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + num_videos_per_prompt = 1 + + # 1. Check inputs + self.check_inputs( + prompt, height, width, negative_prompt, + callback_on_step_end_tensor_inputs, prompt_embeds, negative_prompt_embeds, + ) + self._guidance_scale = guidance_scale + self._attention_kwargs = attention_kwargs + self._interrupt = False + + # 2. Default call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + weight_dtype = self.text_encoder.dtype + + do_classifier_free_guidance = guidance_scale > 1.0 + + # 3. Encode input prompt + prompt_embeds, negative_prompt_embeds = self.encode_prompt( + prompt, negative_prompt, do_classifier_free_guidance, + num_videos_per_prompt=num_videos_per_prompt, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + max_sequence_length=max_sequence_length, + device=device, + ) + if do_classifier_free_guidance: + in_prompt_embeds = negative_prompt_embeds + prompt_embeds + else: + in_prompt_embeds = prompt_embeds + + # 4. Prepare timesteps + if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler): + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, num_inference_steps, device, timesteps, mu=1 + ) + elif isinstance(self.scheduler, FlowUniPCMultistepScheduler): + self.scheduler.set_timesteps(num_inference_steps, device=device, shift=shift) + timesteps = self.scheduler.timesteps + elif isinstance(self.scheduler, FlowDPMSolverMultistepScheduler): + sampling_sigmas = get_sampling_sigmas(num_inference_steps, shift) + timesteps, _ = retrieve_timesteps(self.scheduler, device=device, sigmas=sampling_sigmas) + else: + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, num_inference_steps, device, timesteps + ) + self._num_timesteps = len(timesteps) + + if comfyui_progressbar: + from comfy.utils import ProgressBar + pbar = ProgressBar(num_inference_steps + 2) + + # 5. Prepare latents + if video is not None: + video_length = video.shape[2] + init_video = self.image_processor.preprocess( + rearrange(video, "b c f h w -> (b f) c h w"), height=height, width=width + ) + init_video = init_video.to(dtype=torch.float32) + init_video = rearrange(init_video, "(b f) c h w -> b c f h w", f=video_length) + else: + init_video = None + + latent_channels = self.vae.config.latent_channels + latents = self.prepare_latents( + batch_size * num_videos_per_prompt, + latent_channels, + num_frames, + height, + width, + weight_dtype, + device, + generator, + latents, + ) + + if comfyui_progressbar: + pbar.update(1) + + # 6. Prepare mask latent variables + if init_video is not None: + if (mask_video == 255).all(): + mask_latents = torch.tile( + torch.zeros_like(latents)[:, :1].to(device, weight_dtype), [1, 4, 1, 1, 1] + ) + masked_video_latents = torch.zeros_like(latents).to(device, weight_dtype) + if self.vae.spatial_compression_ratio >= 16: + mask = torch.ones_like(latents).to(device, weight_dtype)[:, :1] + else: + bs, _, video_length, height, width = video.size() + mask_condition = self.mask_processor.preprocess( + rearrange(mask_video, "b c f h w -> (b f) c h w"), height=height, width=width + ) + mask_condition = mask_condition.to(dtype=torch.float32) + mask_condition = rearrange(mask_condition, "(b f) c h w -> b c f h w", f=video_length) + + masked_video = init_video * (torch.tile(mask_condition, [1, 3, 1, 1, 1]) < 0.5) + _, masked_video_latents = self.prepare_mask_latents( + None, masked_video, batch_size, height, width, + weight_dtype, device, generator, do_classifier_free_guidance, noise_aug_strength=None, + ) + + mask_condition = torch.concat([ + torch.repeat_interleave(mask_condition[:, :, 0:1], repeats=4, dim=2), + mask_condition[:, :, 1:] + ], dim=2) + mask_condition = mask_condition.view(bs, mask_condition.shape[2] // 4, 4, height, width) + mask_condition = mask_condition.transpose(1, 2) + mask_latents = resize_mask(1 - mask_condition, masked_video_latents, True).to(device, weight_dtype) + + if self.vae.spatial_compression_ratio >= 16: + mask = TorchF.interpolate( + mask_condition[:, :1], size=latents.size()[-3:], mode='trilinear', align_corners=True + ).to(device, weight_dtype) + if not mask[:, :, 0, :, :].any(): + mask[:, :, 1:, :, :] = 1 + latents = (1 - mask) * masked_video_latents + mask * latents + + if comfyui_progressbar: + pbar.update(1) + + # 7. Prepare camera conditioning + lat_f = (num_frames - 1) // self.vae.temporal_compression_ratio + 1 + lat_h = height // self.vae.spatial_compression_ratio + lat_w = width // self.vae.spatial_compression_ratio + + dit_cond_dict = self.prepare_camera_conditioning( + action_path=action_path, + frame_num=num_frames, + height=height, + width=width, + lat_f=lat_f, + lat_h=lat_h, + lat_w=lat_w, + device=device, + dtype=weight_dtype, + ) + + # 8. Prepare extra step kwargs + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + target_shape = ( + self.vae.latent_channels, + lat_f, + width // self.vae.spatial_compression_ratio, + height // self.vae.spatial_compression_ratio + ) + seq_len = math.ceil( + (target_shape[2] * target_shape[3]) / + (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1] + ) + + # 9. Denoising loop + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self.transformer.num_inference_steps = num_inference_steps + + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + self.transformer.current_steps = i + + if self.interrupt: + continue + + latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents + if hasattr(self.scheduler, "scale_model_input"): + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + if init_video is not None: + mask_input = torch.cat([mask_latents] * 2) if do_classifier_free_guidance else mask_latents + masked_video_latents_input = ( + torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents + ) + y = torch.cat([mask_input, masked_video_latents_input], dim=1).to(device, weight_dtype) + + # Timestep handling + if self.vae.spatial_compression_ratio >= 16 and init_video is not None: + temp_ts = ((mask[0][0][:, ::2, ::2]) * t).flatten() + temp_ts = torch.cat([ + temp_ts, + temp_ts.new_ones(seq_len - temp_ts.size(0)) * t + ]) + temp_ts = temp_ts.unsqueeze(0) + timestep = temp_ts.expand(latent_model_input.shape[0], temp_ts.size(1)) + else: + timestep = t.expand(latent_model_input.shape[0]) + + # Select transformer based on timestep + if self.transformer_2 is not None: + if t >= boundary * self.scheduler.config.num_train_timesteps: + local_transformer = self.transformer_2 + else: + local_transformer = self.transformer + else: + local_transformer = self.transformer + + # Predict noise + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=device): + noise_pred = local_transformer( + x=latent_model_input, + context=in_prompt_embeds, + t=timestep, + seq_len=seq_len, + y=y, + dit_cond_dict=dit_cond_dict, + ) + + # Perform guidance + if do_classifier_free_guidance: + if self.transformer_2 is not None and isinstance(self.guidance_scale, (list, tuple)): + sample_guide_scale = ( + self.guidance_scale[1] + if t >= self.transformer_2.config.boundary * self.scheduler.config.num_train_timesteps + else self.guidance_scale[0] + ) + else: + sample_guide_scale = self.guidance_scale + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + sample_guide_scale * (noise_pred_text - noise_pred_uncond) + + # Compute previous noisy sample + latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + + if self.vae.spatial_compression_ratio >= 16 and not mask[:, :, 0, :, :].any(): + latents = (1 - mask) * masked_video_latents + mask * latents + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) + + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + if comfyui_progressbar: + pbar.update(1) + + # 10. Decode latents + if output_type == "numpy": + video = self.decode_latents(latents) + elif not output_type == "latent": + video = self.decode_latents(latents) + video = self.video_processor.postprocess_video(video=video, output_type=output_type) + else: + video = latents + + # Offload models + self.maybe_free_model_hooks() + + if not return_dict: + video = torch.from_numpy(video) + + return LingBotWorldPipelineOutput(videos=video) diff --git a/videox_fun/utils/cam_utils.py b/videox_fun/utils/cam_utils.py new file mode 100644 index 0000000..a131dd7 --- /dev/null +++ b/videox_fun/utils/cam_utils.py @@ -0,0 +1,231 @@ +"""Borrowed from LingBot-World/wan/utils/cam_utils.py. +""" + +import numpy as np +import torch +from scipy.interpolate import interp1d +from scipy.spatial.transform import Rotation, Slerp + + +def interpolate_camera_poses( + src_indices: np.ndarray, + src_rot_mat: np.ndarray, + src_trans_vec: np.ndarray, + tgt_indices: np.ndarray, +) -> torch.Tensor: + """ + Interpolate camera poses using linear interpolation for translation + and Slerp for rotation. + + Args: + src_indices: Source frame indices + src_rot_mat: Source rotation matrices, shape [N, 3, 3] + src_trans_vec: Source translation vectors, shape [N, 3] + tgt_indices: Target frame indices for interpolation + + Returns: + Interpolated camera poses as 4x4 matrices, shape [M, 4, 4] + """ + # interpolate translation + interp_func_trans = interp1d( + src_indices, + src_trans_vec, + axis=0, + kind='linear', + bounds_error=False, + fill_value="extrapolate", + ) + interpolated_trans_vec = interp_func_trans(tgt_indices) + + # interpolate rotation + src_quat_vec = Rotation.from_matrix(src_rot_mat) + # ensure there is no sudden change in qw + quats = src_quat_vec.as_quat().copy() # [N, 4] + for i in range(1, len(quats)): + if np.dot(quats[i], quats[i-1]) < 0: + quats[i] = -quats[i] + src_quat_vec = Rotation.from_quat(quats) + slerp_func_rot = Slerp(src_indices, src_quat_vec) + interpolated_rot_quat = slerp_func_rot(tgt_indices) + interpolated_rot_mat = interpolated_rot_quat.as_matrix() + + poses = np.zeros((len(tgt_indices), 4, 4)) + poses[:, :3, :3] = interpolated_rot_mat + poses[:, :3, 3] = interpolated_trans_vec + poses[:, 3, 3] = 1.0 + return torch.from_numpy(poses).float() + + +def SE3_inverse(T: torch.Tensor) -> torch.Tensor: + """ + Compute the inverse of SE3 transformation matrices. + + Args: + T: SE3 transformation matrices, shape [B, 4, 4] + + Returns: + Inverse transformation matrices, shape [B, 4, 4] + """ + Rot = T[:, :3, :3] # [B, 3, 3] + trans = T[:, :3, 3:] # [B, 3, 1] + R_inv = Rot.transpose(-1, -2) + t_inv = -torch.bmm(R_inv, trans) + T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1) + T_inv[:, :3, :3] = R_inv + T_inv[:, :3, 3:] = t_inv + return T_inv + + +def compute_relative_poses( + c2ws_mat: torch.Tensor, + framewise: bool = False, + normalize_trans: bool = True, +) -> torch.Tensor: + """ + Compute relative poses from camera-to-world matrices. + + Args: + c2ws_mat: Camera-to-world matrices, shape [F, 4, 4] + framewise: If True, compute frame-to-frame relative poses + normalize_trans: If True, normalize translations to unit max norm + + Returns: + Relative pose matrices, shape [F, 4, 4] + """ + ref_w2cs = SE3_inverse(c2ws_mat[0:1]) + relative_poses = torch.matmul(ref_w2cs, c2ws_mat) + # ensure identity matrix for 1st frame + relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype) + if framewise: + # compute pose between i and i+1 + relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:]) + relative_poses[1:] = relative_poses_framewise + if normalize_trans: + # note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning." + translations = relative_poses[:, :3, 3] # [f, 3] + max_norm = torch.norm(translations, dim=-1).max() + # only normalize when moving + if max_norm > 0: + relative_poses[:, :3, 3] = translations / max_norm + return relative_poses + + +@torch.no_grad() +def create_meshgrid( + n_frames: int, + height: int, + width: int, + bias: float = 0.5, + device='cuda', + dtype=torch.float32 +) -> torch.Tensor: + """ + Create a meshgrid for ray generation. + + Args: + n_frames: Number of frames + height: Image height + width: Image width + bias: Pixel center bias (0.5 for pixel center) + device: Torch device + dtype: Torch dtype + + Returns: + Grid coordinates, shape [F, H*W, 2] + """ + x_range = torch.arange(width, device=device, dtype=dtype) + y_range = torch.arange(height, device=device, dtype=dtype) + grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij') + grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2] + grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2] + return grid_xy + + +def get_plucker_embeddings( + c2ws_mat: torch.Tensor, + Ks: torch.Tensor, + height: int, + width: int, +) -> torch.Tensor: + """ + Generate Plucker ray embeddings from camera parameters. + + Args: + c2ws_mat: Camera-to-world matrices, shape [F, 4, 4] + Ks: Camera intrinsics [fx, fy, cx, cy], shape [F, 4] + height: Image height + width: Image width + + Returns: + Plucker embeddings (ray_origin + ray_direction), shape [F, H, W, 6] + """ + n_frames = c2ws_mat.shape[0] + grid_xy = create_meshgrid(n_frames, height, width, device=c2ws_mat.device, dtype=c2ws_mat.dtype) # [f, h*w, 2] + fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1] + + i = grid_xy[..., 0] # [f, h*w] + j = grid_xy[..., 1] # [f, h*w] + zs = torch.ones_like(i) # [f, h*w] + xs = (i - cx) / fx * zs + ys = (j - cy) / fy * zs + + directions = torch.stack([xs, ys, zs], dim=-1) # [f, h*w, 3] + directions = directions / directions.norm(dim=-1, keepdim=True) # [f, h*w, 3] + + rays_d = directions @ c2ws_mat[:, :3, :3].transpose(-1, -2) # [f, h*w, 3] + rays_o = c2ws_mat[:, :3, 3] # [f, 3] + rays_o = rays_o[:, None, :].expand_as(rays_d) # [f, h*w, 3] + + # Plucker coordinates: origin + direction (refer to apt2) + plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1) # [f, h*w, 6] + plucker_embeddings = plucker_embeddings.view([n_frames, height, width, 6]) # [f, h, w, 6] + return plucker_embeddings + + +def get_Ks_transformed( + Ks: torch.Tensor, + height_org: int, + width_org: int, + height_resize: int, + width_resize: int, + height_final: int, + width_final: int, +) -> torch.Tensor: + """ + Transform camera intrinsics for image resizing and cropping. + + Args: + Ks: Original camera intrinsics [fx, fy, cx, cy], shape [F, 4] + height_org: Original image height + width_org: Original image width + height_resize: Resized image height + width_resize: Resized image width + height_final: Final cropped image height + width_final: Final cropped image width + + Returns: + Transformed camera intrinsics, shape [F, 4] + """ + fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1] + + scale_x = width_resize / width_org + scale_y = height_resize / height_org + + fx_resize = fx * scale_x + fy_resize = fy * scale_y + cx_resize = cx * scale_x + cy_resize = cy * scale_y + + crop_offset_x = (width_resize - width_final) / 2 + crop_offset_y = (height_resize - height_final) / 2 + + cx_final = cx_resize - crop_offset_x + cy_final = cy_resize - crop_offset_y + + Ks_transformed = torch.zeros_like(Ks) + Ks_transformed[:, 0:1] = fx_resize + Ks_transformed[:, 1:2] = fy_resize + Ks_transformed[:, 2:3] = cx_final + Ks_transformed[:, 3:4] = cy_final + + return Ks_transformed